Add offline unit test suite for chat_backend (#12)
Unit Tests / test (push) Successful in 9s
Unit Tests / test (push) Successful in 9s
Closes #5 ## Summary - Replaces the three scattered test modules (`chat_backend/tests.py`, `services/tests.py`, `services/prompt_classifier/tests.py`) with a `chat_backend/tests/` package: **242 deterministic tests plus 6 opt-in live-Ollama checks**, up from 10 tests (3 of which were skipped and 4 of which were never even discovered). - The suite runs fully offline — no Ollama, Chroma, SMTP or network access. LangChain runnables are replaced by a small `FakeChain`, Chroma/embeddings are mocked, email uses Django's locmem backend, and blobs go through `DatabaseStorage`. - New `llm_be/test_runner.py` (wired via `TEST_RUNNER`) sets `SKIP_RAG_INIT=1` and an MD5 password hasher, so the suite cannot accidentally reach a model server and finishes in ~5s on SQLite (~17s on Postgres) instead of ~30s. ## Coverage | Area | File | |------|------| | `TimeInfoBase.save`, slugs, `get_duration`, `file_exists`, cascades | `test_models.py` | | `DatabaseStorage` save/open/exists/size/listdir/delete/times | `test_storage.py` | | JWT claim, prompt/user/feedback/document serializers | `test_serializers.py` | | auth + token, invite, feedback, company users, set-password, TOS | `test_views_users.py` | | conversation list/order/create/detail/soft-delete | `test_views_conversations.py` | | all four analytics endpoints, including empty-month behaviour | `test_views_analytics.py` | | workspace + document upload/list/detail, 404 and 400 paths | `test_views_documents.py` | | prompt classifier rules/parsing, moderation fail-safe, title cleanup | `test_services_classifiers.py` | | CSV/XLSX/DOCX/PDF analysis, plot generation, error payloads | `test_services_data_analysis.py` | | loader selection, filename sanitising, ingest, temp-file cleanup, search filters | `test_services_rag.py` | | history formatting and streaming | `test_services_llm.py` | | document re-index on create/delete, `SKIP_RAG_INIT` guard | `test_signals.py` | | consumer DB helpers, LangGraph nodes (moderation, classification, generation, search flags), websocket routes | `test_consumers.py` | Live checks (non-deterministic, need a model server): ```bash cd llm_be RUN_LIVE_OLLAMA_TESTS=1 uv run python manage.py test chat_backend.tests.test_live_ollama ``` ## Bugs the tests surfaced (fixed here) 1. **`ConversationDetailView.post` silently dropped every prompt.** `import datetime` shadowed `from datetime import datetime`, so `datetime.now()` raised `AttributeError` inside a bare `except` and the endpoint returned 200 without saving. Now uses `timezone.now()`. 2. **Prompt attachments never reached the LLM.** `get_conversation_file_async` (both consumers) did `sync_to_async(prompt.file.read)` — with `DatabaseStorage` the attribute access itself opens the blob, i.e. a DB query in async context, raising `SynchronousOnlyOperation` that was swallowed and returned `(None, None)`. The read now happens inside the thread. 3. **`DatabaseStorage._save` crashed on a str-backed `ContentFile`** (`TypeError: sequence item 0: expected a bytes-like object`); chunks are encoded when needed. 4. **`services/prompt_classifier/__init__.,py`** (note the comma) meant the directory was only an implicit namespace package, which is why its test module was never collected. Renamed, and its duplicate live-Ollama tests folded into `test_live_ollama.py`. Known-broken paths deliberately left untested and unchanged: `reset_password` / `ResetUserPassword` reference an unimported `requests` plus undefined locals, and `DocumentDetailView.get` references an undefined `workspaces` on its success path. Worth a follow-up ticket. ## Test plan - [x] `cd llm_be && uv run python manage.py test` → 248 tests, OK (6 skipped, all opt-in live) - [x] Same suite against Postgres 16 (`DATABASE_URL=postgres://…`) → OK, matching the containerized run in `deploy.yml` - [x] `uv run black` clean on all added files - [ ] Gitea Actions **Unit Tests** + **CI** green on this PRReviewed-on: #12
This commit was merged in pull request #12.
This commit is contained in:
@@ -196,8 +196,8 @@ async def get_conversation_file_async(conversation_id):
|
||||
).exclude(file='').order_by('created').afirst()
|
||||
|
||||
if prompt_with_file and prompt_with_file.file:
|
||||
# You must use sync_to_async to access the file's binary content
|
||||
file_data = await sync_to_async(prompt_with_file.file.read)()
|
||||
# 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:
|
||||
|
||||
@@ -7,6 +7,7 @@ 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
|
||||
@@ -139,7 +140,8 @@ async def get_conversation_file_async(conversation_id):
|
||||
).exclude(file='').order_by('created').afirst()
|
||||
|
||||
if prompt_with_file and prompt_with_file.file:
|
||||
file_data = await sync_to_async(prompt_with_file.file.read)()
|
||||
# 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:
|
||||
|
||||
@@ -1,31 +0,0 @@
|
||||
import os
|
||||
from unittest import TestCase, mock
|
||||
from unittest.mock import MagicMock, patch, AsyncMock
|
||||
from typing import List, Dict, Any
|
||||
|
||||
from django.test import TestCase as DjangoTestCase
|
||||
|
||||
from chat_backend.services.rag_services import (
|
||||
RAGService,
|
||||
SyncRAGService,
|
||||
AsyncRAGService,
|
||||
)
|
||||
from chat_backend.models import Conversation, Prompt, DocumentWorkspace, Document
|
||||
from chat_backend.services.prompt_classifier.prompt_classifier import PromptClassifier, PromptType
|
||||
from parameterized import parameterized
|
||||
|
||||
|
||||
|
||||
class PromptClassifierTestCase(TestCase):
|
||||
def setUp(self):
|
||||
self.service = PromptClassifier()
|
||||
|
||||
@parameterized.expand([
|
||||
["Tell me a joke",PromptType.GENERAL_CHAT],
|
||||
["Create an image of a dog for me",PromptType.IMAGE_GENERATION],
|
||||
["highlight the features of the backyard playset if they were to choose us and make the language more long form",PromptType.GENERAL_CHAT],
|
||||
["Great, can you make it about a duck now", PromptType.IMAGE_GENERATION],
|
||||
])
|
||||
def test_prompt_classification(self, prompt, expected_output):
|
||||
result = self.service.classify(prompt)
|
||||
self.assertEqual(result, expected_output)
|
||||
@@ -1,32 +0,0 @@
|
||||
import os
|
||||
import unittest
|
||||
from unittest import TestCase
|
||||
|
||||
from chat_backend.services.prompt_classifier.prompt_classifier import (
|
||||
PromptClassifier,
|
||||
PromptType,
|
||||
)
|
||||
from parameterized import parameterized
|
||||
|
||||
|
||||
@unittest.skipIf(
|
||||
os.environ.get("SKIP_RAG_INIT", "").lower() in {"1", "true", "yes"},
|
||||
"Requires live Ollama; skipped when SKIP_RAG_INIT is set",
|
||||
)
|
||||
class PromptClassifierTestCase(TestCase):
|
||||
def setUp(self):
|
||||
self.service = PromptClassifier()
|
||||
|
||||
@parameterized.expand(
|
||||
[
|
||||
["Tell me a joke", PromptType.GENERAL_CHAT],
|
||||
["Create an image of a dog for me", PromptType.IMAGE_GENERATION],
|
||||
[
|
||||
"highlight the features of the backyard playset if they were to choose us and make the language more long form",
|
||||
PromptType.GENERAL_CHAT,
|
||||
],
|
||||
]
|
||||
)
|
||||
def test_prompt_classification(self, prompt, expected_output):
|
||||
result = self.service.classify(prompt)
|
||||
self.assertEqual(result, expected_output)
|
||||
@@ -27,13 +27,17 @@ class DatabaseStorage(Storage):
|
||||
def _save(self, name, content):
|
||||
name = self.get_available_name(name)
|
||||
if hasattr(content, "chunks"):
|
||||
data = b"".join(chunk for chunk in content.chunks())
|
||||
chunks = content.chunks()
|
||||
else:
|
||||
data = content.read()
|
||||
if isinstance(data, str):
|
||||
data = data.encode("utf-8")
|
||||
chunks = [content.read()]
|
||||
data = b"".join(
|
||||
chunk.encode("utf-8") if isinstance(chunk, str) else chunk
|
||||
for chunk in chunks
|
||||
)
|
||||
|
||||
content_type = getattr(content, "content_type", None) or mimetypes.guess_type(name)[0]
|
||||
content_type = (
|
||||
getattr(content, "content_type", None) or mimetypes.guess_type(name)[0]
|
||||
)
|
||||
StoredFile = self._model()
|
||||
with transaction.atomic():
|
||||
StoredFile.objects.update_or_create(
|
||||
@@ -56,8 +60,10 @@ class DatabaseStorage(Storage):
|
||||
prefix = path.rstrip("/")
|
||||
if prefix:
|
||||
prefix = f"{prefix}/"
|
||||
names = self._model().objects.filter(name__startswith=prefix).values_list(
|
||||
"name", flat=True
|
||||
names = (
|
||||
self._model()
|
||||
.objects.filter(name__startswith=prefix)
|
||||
.values_list("name", flat=True)
|
||||
)
|
||||
dirs: set[str] = set()
|
||||
files: list[str] = []
|
||||
@@ -88,4 +94,6 @@ class DatabaseStorage(Storage):
|
||||
return self._model().objects.values_list("created", flat=True).get(name=name)
|
||||
|
||||
def get_modified_time(self, name):
|
||||
return self._model().objects.values_list("last_modified", flat=True).get(name=name)
|
||||
return (
|
||||
self._model().objects.values_list("last_modified", flat=True).get(name=name)
|
||||
)
|
||||
|
||||
@@ -1,194 +0,0 @@
|
||||
from django.test import TestCase
|
||||
|
||||
# Create your tests here.
|
||||
from django.test import TestCase, Client
|
||||
from django.urls import reverse
|
||||
from django.contrib.auth.models import User
|
||||
from rest_framework.test import APIClient, APITestCase
|
||||
from rest_framework import status
|
||||
from .models import DocumentWorkspace, Document, Company
|
||||
from django.contrib.auth import get_user_model
|
||||
import tempfile
|
||||
from django.core.files.uploadedfile import SimpleUploadedFile
|
||||
from parameterized import parameterized
|
||||
|
||||
|
||||
# Minimal valid PDF bytes
|
||||
VALID_PDF_BYTES = (
|
||||
b"%PDF-1.3\n"
|
||||
b"1 0 obj\n"
|
||||
b"<< /Type /Catalog /Pages 2 0 R >>\n"
|
||||
b"endobj\n"
|
||||
b"2 0 obj\n"
|
||||
b"<< /Type /Pages /Kids [3 0 R] /Count 1 >>\n"
|
||||
b"endobj\n"
|
||||
b"3 0 obj\n"
|
||||
b"<< /Type /Page /Parent 2 0 R /Resources << >> /MediaBox [0 0 612 792] /Contents 4 0 R >>\n"
|
||||
b"endobj\n"
|
||||
b"4 0 obj\n"
|
||||
b"<< /Length 44 >>\n"
|
||||
b"stream\n"
|
||||
b"BT /F1 12 Tf 72 720 Td (Test PDF) Tj ET\n"
|
||||
b"endstream\n"
|
||||
b"endobj\n"
|
||||
b"xref\n"
|
||||
b"0 5\n"
|
||||
b"0000000000 65535 f \n"
|
||||
b"0000000009 00000 n \n"
|
||||
b"0000000058 00000 n \n"
|
||||
b"0000000117 00000 n \n"
|
||||
b"0000000223 00000 n \n"
|
||||
b"trailer\n"
|
||||
b"<< /Size 5 /Root 1 0 R >>\n"
|
||||
b"startxref\n"
|
||||
b"317\n"
|
||||
b"%%EOF"
|
||||
)
|
||||
|
||||
|
||||
class DocumentWorkspaceViewsTestCase(APITestCase):
|
||||
def setUp(self):
|
||||
self.company = Company.objects.create(
|
||||
name="test", state="IL", zipcode="60189", address="1968 Greensboro Dr"
|
||||
)
|
||||
self.user = get_user_model().objects.create_user(
|
||||
company=self.company,
|
||||
username="testuser",
|
||||
password="testpass123",
|
||||
email="test@test.com",
|
||||
)
|
||||
|
||||
self.client = APIClient()
|
||||
self.client.force_authenticate(user=self.user)
|
||||
|
||||
self.workspace = DocumentWorkspace.objects.create(
|
||||
company=self.user.company, name="Test Workspace"
|
||||
)
|
||||
|
||||
def test_list_workspaces(self):
|
||||
url = reverse("document_workspaces")
|
||||
response = self.client.get(url)
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
self.assertEqual(len(response.data), 1)
|
||||
self.assertEqual(response.data[0]["name"], "Test Workspace")
|
||||
|
||||
def test_create_workspace(self):
|
||||
url = reverse("document_workspaces")
|
||||
data = {"name": "New Workspace"}
|
||||
response = self.client.post(url, data, format="json")
|
||||
self.assertEqual(response.status_code, status.HTTP_201_CREATED)
|
||||
self.assertEqual(DocumentWorkspace.objects.count(), 2)
|
||||
|
||||
def test_retrieve_workspace(self):
|
||||
url = reverse("document_workspaces")
|
||||
response = self.client.get(url)
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
self.assertEqual(response.data[0]["name"], "Test Workspace")
|
||||
|
||||
# def test_update_workspace(self):
|
||||
# url = reverse('document_workspaces')
|
||||
# data = {
|
||||
# 'name': 'Updated Workspace'
|
||||
# }
|
||||
# response = self.client.post(url, data, format='json')
|
||||
# self.assertEqual(response.status_code, status.HTTP_201_CREATED)
|
||||
# self.workspace.refresh_from_db()
|
||||
# self.assertEqual(self.workspace.name, 'Updated Workspace')
|
||||
|
||||
# def test_delete_workspace(self):
|
||||
# url = reverse('document_workspaces', args=[self.workspace.id])
|
||||
# response = self.client.delete(url)
|
||||
# self.assertEqual(response.status_code, status.HTTP_204_NO_CONTENT)
|
||||
# self.assertEqual(DocumentWorkspace.objects.count(), 0)
|
||||
|
||||
|
||||
class DocumentViewsTestCase(APITestCase):
|
||||
def setUp(self):
|
||||
self.company = Company.objects.create(
|
||||
name="test", state="IL", zipcode="60189", address="1968 Greensboro Dr"
|
||||
)
|
||||
self.user = get_user_model().objects.create_user(
|
||||
company=self.company,
|
||||
username="testuser",
|
||||
password="testpass123",
|
||||
email="test@test.com",
|
||||
)
|
||||
|
||||
self.client = APIClient()
|
||||
self.client.force_authenticate(user=self.user)
|
||||
|
||||
self.workspace = DocumentWorkspace.objects.create(
|
||||
company=self.user.company, name="Test Workspace"
|
||||
)
|
||||
|
||||
# Create a test file
|
||||
self.test_file = SimpleUploadedFile(
|
||||
"test.pdf", VALID_PDF_BYTES, content_type="application/pdf"
|
||||
)
|
||||
|
||||
def test_upload_document(self):
|
||||
url = reverse("documents")
|
||||
data = {"file": self.test_file}
|
||||
response = self.client.post(url, data, format="multipart")
|
||||
self.assertEqual(response.status_code, status.HTTP_201_CREATED)
|
||||
self.assertEqual(Document.objects.count(), 1)
|
||||
|
||||
document = Document.objects.first()
|
||||
self.assertEqual(document.workspace.id, self.workspace.id)
|
||||
self.assertTrue(document.processed) # Should be False initially
|
||||
|
||||
def test_list_documents(self):
|
||||
# First create a document
|
||||
Document.objects.create(workspace=self.workspace, file=self.test_file)
|
||||
|
||||
url = reverse("documents")
|
||||
response = self.client.get(url)
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
self.assertEqual(len(response.data), 1)
|
||||
self.assertIn("test", response.data[0]["file"])
|
||||
self.assertIn("pdf", response.data[0]["file"])
|
||||
|
||||
# def test_delete_document(self):
|
||||
# document = Document.objects.create(
|
||||
# workspace=self.workspace,
|
||||
# file=self.test_file
|
||||
# )
|
||||
|
||||
# url = reverse('document-detail', args=[document.id])
|
||||
# response = self.client.delete(url)
|
||||
# self.assertEqual(response.status_code, status.HTTP_204_NO_CONTENT)
|
||||
# self.assertEqual(Document.objects.count(), 0)
|
||||
|
||||
def test_upload_invalid_file(self):
|
||||
url = reverse("documents")
|
||||
data = {"file": "not a file"}
|
||||
response = self.client.post(url, data, format="multipart")
|
||||
self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST)
|
||||
|
||||
def test_access_other_users_documents(self):
|
||||
# Create another user
|
||||
other_company = Company.objects.create(
|
||||
name="test2", state="IL", zipcode="60189", address="1968 Greensboro Dr"
|
||||
)
|
||||
other_user = get_user_model().objects.create_user(
|
||||
company=other_company,
|
||||
username="otheruser",
|
||||
password="otherpass123",
|
||||
email="testing2@test.com",
|
||||
)
|
||||
other_workspace = DocumentWorkspace.objects.create(
|
||||
company=other_user.company, name="Other Workspace"
|
||||
)
|
||||
other_document = Document.objects.create(
|
||||
workspace=other_workspace, file=self.test_file
|
||||
)
|
||||
|
||||
# Try to access the other user's document
|
||||
url = reverse("documents_details", kwargs={"document_id": other_document.id})
|
||||
response = self.client.get(url)
|
||||
self.assertEqual(response.status_code, status.HTTP_404_NOT_FOUND)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,6 @@
|
||||
"""Unit tests for the chat_backend app.
|
||||
|
||||
Every test in this package runs offline: no Ollama, Chroma, SMTP or network calls.
|
||||
LLM chains are replaced with fakes, email uses the locmem backend Django installs
|
||||
for tests, and file storage is the database-backed ``DatabaseStorage``.
|
||||
"""
|
||||
@@ -0,0 +1,131 @@
|
||||
"""Shared fixtures/helpers for chat_backend tests."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from django.contrib.auth import get_user_model
|
||||
from django.core.files.uploadedfile import SimpleUploadedFile
|
||||
|
||||
from chat_backend.models import (
|
||||
Company,
|
||||
Conversation,
|
||||
Document,
|
||||
DocumentWorkspace,
|
||||
Prompt,
|
||||
)
|
||||
|
||||
# Minimal valid PDF bytes — enough for pypdf/PyPDFLoader to parse.
|
||||
VALID_PDF_BYTES = (
|
||||
b"%PDF-1.3\n"
|
||||
b"1 0 obj\n"
|
||||
b"<< /Type /Catalog /Pages 2 0 R >>\n"
|
||||
b"endobj\n"
|
||||
b"2 0 obj\n"
|
||||
b"<< /Type /Pages /Kids [3 0 R] /Count 1 >>\n"
|
||||
b"endobj\n"
|
||||
b"3 0 obj\n"
|
||||
b"<< /Type /Page /Parent 2 0 R /Resources << >> /MediaBox [0 0 612 792] /Contents 4 0 R >>\n"
|
||||
b"endobj\n"
|
||||
b"4 0 obj\n"
|
||||
b"<< /Length 44 >>\n"
|
||||
b"stream\n"
|
||||
b"BT /F1 12 Tf 72 720 Td (Test PDF) Tj ET\n"
|
||||
b"endstream\n"
|
||||
b"endobj\n"
|
||||
b"xref\n"
|
||||
b"0 5\n"
|
||||
b"0000000000 65535 f \n"
|
||||
b"0000000009 00000 n \n"
|
||||
b"0000000058 00000 n \n"
|
||||
b"0000000117 00000 n \n"
|
||||
b"0000000223 00000 n \n"
|
||||
b"trailer\n"
|
||||
b"<< /Size 5 /Root 1 0 R >>\n"
|
||||
b"startxref\n"
|
||||
b"317\n"
|
||||
b"%%EOF"
|
||||
)
|
||||
|
||||
CSV_BYTES = b"name,sales,units\nalice,100,3\nbob,200,5\ncarol,300,9\n"
|
||||
|
||||
|
||||
def pdf_bytes_with_text(text: str = "Quarterly Report") -> bytes:
|
||||
"""A real, text-extractable PDF (matplotlib is already a dependency)."""
|
||||
import io
|
||||
|
||||
import matplotlib
|
||||
|
||||
matplotlib.use("Agg")
|
||||
import matplotlib.pyplot as plt
|
||||
|
||||
figure = plt.figure()
|
||||
figure.text(0.1, 0.5, text)
|
||||
buffer = io.BytesIO()
|
||||
figure.savefig(buffer, format="pdf")
|
||||
plt.close(figure)
|
||||
return buffer.getvalue()
|
||||
|
||||
|
||||
def docx_bytes_with_text(*paragraphs: str) -> bytes:
|
||||
import io
|
||||
|
||||
import docx
|
||||
|
||||
document = docx.Document()
|
||||
for paragraph in paragraphs:
|
||||
document.add_paragraph(paragraph)
|
||||
buffer = io.BytesIO()
|
||||
document.save(buffer)
|
||||
return buffer.getvalue()
|
||||
|
||||
|
||||
def make_company(name: str = "Acme", **kwargs) -> Company:
|
||||
defaults = {
|
||||
"state": "IL",
|
||||
"zipcode": "60189",
|
||||
"address": "1968 Greensboro Dr",
|
||||
}
|
||||
defaults.update(kwargs)
|
||||
return Company.objects.create(name=name, **defaults)
|
||||
|
||||
|
||||
def make_user(
|
||||
email: str = "test@test.com",
|
||||
password: str | None = "testpass123",
|
||||
company: Company | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
user_model = get_user_model()
|
||||
return user_model.objects.create_user(
|
||||
username=kwargs.pop("username", email),
|
||||
email=email,
|
||||
password=password,
|
||||
company=company,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
def make_conversation(user=None, title: str = "Test Conversation", **kwargs):
|
||||
return Conversation.objects.create(user=user, title=title, **kwargs)
|
||||
|
||||
|
||||
def make_prompt(
|
||||
conversation, message: str = "hello", user_created: bool = True, **kwargs
|
||||
):
|
||||
return Prompt.objects.create(
|
||||
conversation=conversation,
|
||||
message=message,
|
||||
user_created=user_created,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
def make_workspace(company, name: str = "Test Workspace") -> DocumentWorkspace:
|
||||
return DocumentWorkspace.objects.create(company=company, name=name)
|
||||
|
||||
|
||||
def pdf_upload(name: str = "test.pdf") -> SimpleUploadedFile:
|
||||
return SimpleUploadedFile(name, VALID_PDF_BYTES, content_type="application/pdf")
|
||||
|
||||
|
||||
def make_document(workspace, name: str = "test.pdf") -> Document:
|
||||
return Document.objects.create(workspace=workspace, file=pdf_upload(name))
|
||||
@@ -0,0 +1,44 @@
|
||||
"""Stand-ins for LangChain runnables so tests never reach Ollama."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
class FakeChain:
|
||||
"""Mimics the subset of the Runnable API the services use.
|
||||
|
||||
Records every payload it was invoked with, so tests can assert on the
|
||||
prompt variables a service builds.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
response: str = "",
|
||||
chunks: list[str] | None = None,
|
||||
error: Exception | None = None,
|
||||
):
|
||||
self.response = response
|
||||
self.chunks = chunks if chunks is not None else [response]
|
||||
self.error = error
|
||||
self.calls: list[dict] = []
|
||||
|
||||
def invoke(self, payload):
|
||||
self.calls.append(payload)
|
||||
if self.error:
|
||||
raise self.error
|
||||
return self.response
|
||||
|
||||
async def ainvoke(self, payload):
|
||||
return self.invoke(payload)
|
||||
|
||||
def stream(self, payload):
|
||||
self.calls.append(payload)
|
||||
if self.error:
|
||||
raise self.error
|
||||
yield from self.chunks
|
||||
|
||||
async def astream(self, payload):
|
||||
self.calls.append(payload)
|
||||
if self.error:
|
||||
raise self.error
|
||||
for chunk in self.chunks:
|
||||
yield chunk
|
||||
@@ -0,0 +1,373 @@
|
||||
from unittest import mock
|
||||
|
||||
from asgiref.sync import sync_to_async
|
||||
from channels.testing import WebsocketCommunicator
|
||||
from django.test import TestCase, TransactionTestCase, override_settings
|
||||
from langchain_core.messages import AIMessage, HumanMessage
|
||||
from parameterized import parameterized
|
||||
|
||||
from chat_backend import consumers, consumers_graph
|
||||
from chat_backend.models import Conversation, Prompt, PromptMetric
|
||||
from chat_backend.routing import websocket_urlpatterns
|
||||
from chat_backend.services.moderation_classifier import ModerationLabel
|
||||
from chat_backend.services.prompt_classifier.prompt_classifier import PromptType
|
||||
from llm_be.asgi import application
|
||||
|
||||
from .factories import make_company, make_conversation, make_user, make_workspace
|
||||
|
||||
|
||||
class DatabaseHelperTestCase(TransactionTestCase):
|
||||
"""The ``database_sync_to_async`` helpers both consumers share.
|
||||
|
||||
``TransactionTestCase`` is required: ``database_sync_to_async`` closes old
|
||||
connections, which a class-wide atomic block (plain ``TestCase``) cannot survive.
|
||||
"""
|
||||
|
||||
def setUp(self):
|
||||
self.company = make_company()
|
||||
self.user = make_user(company=self.company)
|
||||
self.workspace = make_workspace(self.company)
|
||||
self.conversation = make_conversation(user=self.user)
|
||||
|
||||
@parameterized.expand([("websocket", consumers), ("langgraph", consumers_graph)])
|
||||
async def test_create_conversation(self, _name, module):
|
||||
conversation_id = await module.create_conversation(
|
||||
"first prompt", self.user.email, "Weather Inquiry"
|
||||
)
|
||||
|
||||
stored = await sync_to_async(
|
||||
lambda: list(
|
||||
Conversation.objects.filter(
|
||||
id=conversation_id, user=self.user
|
||||
).values_list("title", flat=True)
|
||||
)
|
||||
)()
|
||||
self.assertEqual(stored, ["Weather Inquiry"])
|
||||
|
||||
@parameterized.expand([("websocket", consumers), ("langgraph", consumers_graph)])
|
||||
async def test_get_workspace(self, _name, module):
|
||||
workspace = await module.get_workspace(self.conversation.id)
|
||||
|
||||
self.assertEqual(workspace.id, self.workspace.id)
|
||||
|
||||
@parameterized.expand([("websocket", consumers), ("langgraph", consumers_graph)])
|
||||
async def test_get_messages_stores_prompt_and_returns_history(self, _name, module):
|
||||
messages, prompt_instance = await module.get_messages(
|
||||
self.conversation.id, "what is the weather"
|
||||
)
|
||||
|
||||
self.assertEqual([type(message) for message in messages], [HumanMessage])
|
||||
self.assertEqual(messages[0].content, "what is the weather")
|
||||
self.assertEqual(prompt_instance.conversation_id, self.conversation.id)
|
||||
self.assertTrue(prompt_instance.user_created)
|
||||
|
||||
@parameterized.expand([("websocket", consumers), ("langgraph", consumers_graph)])
|
||||
async def test_get_messages_maps_assistant_turns(self, _name, module):
|
||||
await sync_to_async(Prompt.objects.create)(
|
||||
conversation=self.conversation, message="earlier answer", user_created=False
|
||||
)
|
||||
|
||||
messages, _prompt = await module.get_messages(self.conversation.id, "follow up")
|
||||
|
||||
self.assertEqual(
|
||||
[type(message) for message in messages], [AIMessage, HumanMessage]
|
||||
)
|
||||
|
||||
async def test_get_messages_stores_attachment_and_inlines_csv(self):
|
||||
messages, prompt_instance = await consumers.get_messages(
|
||||
self.conversation.id, "analyze this", b"name,sales\nalice,10\n", "csv"
|
||||
)
|
||||
|
||||
self.assertEqual(prompt_instance.file_type, "csv")
|
||||
self.assertIn("The file type is csv", messages[-1].content)
|
||||
self.assertIn("alice,10", messages[-1].content)
|
||||
|
||||
@parameterized.expand([("websocket", consumers), ("langgraph", consumers_graph)])
|
||||
async def test_get_conversation_file_returns_stored_attachment(self, _name, module):
|
||||
await consumers.get_messages(
|
||||
self.conversation.id, "analyze this", b"name,sales\nalice,10\n", "csv"
|
||||
)
|
||||
|
||||
file_data, file_type = await module.get_conversation_file_async(
|
||||
self.conversation.id
|
||||
)
|
||||
|
||||
self.assertEqual(file_data, b"name,sales\nalice,10\n")
|
||||
self.assertEqual(file_type, "csv")
|
||||
|
||||
@parameterized.expand([("websocket", consumers), ("langgraph", consumers_graph)])
|
||||
async def test_get_conversation_file_without_attachment(self, _name, module):
|
||||
file_data, file_type = await module.get_conversation_file_async(
|
||||
self.conversation.id
|
||||
)
|
||||
|
||||
self.assertIsNone(file_data)
|
||||
self.assertIsNone(file_type)
|
||||
|
||||
@parameterized.expand([("websocket", consumers), ("langgraph", consumers_graph)])
|
||||
async def test_save_generated_message(self, _name, module):
|
||||
await module.save_generated_message(self.conversation.id, "generated answer")
|
||||
|
||||
stored = await sync_to_async(
|
||||
lambda: list(
|
||||
Prompt.objects.filter(conversation=self.conversation).values(
|
||||
"message", "user_created"
|
||||
)
|
||||
)
|
||||
)()
|
||||
self.assertEqual(
|
||||
stored, [{"message": "generated answer", "user_created": False}]
|
||||
)
|
||||
|
||||
@parameterized.expand([("websocket", consumers), ("langgraph", consumers_graph)])
|
||||
async def test_prompt_metric_lifecycle(self, _name, module):
|
||||
metric = await module.create_prompt_metric(
|
||||
prompt_id=7,
|
||||
prompt="what is the weather",
|
||||
has_file=False,
|
||||
file_type="",
|
||||
model_name="llama3.2",
|
||||
conversation_id=self.conversation.id,
|
||||
)
|
||||
|
||||
self.assertEqual(metric.prompt_length, len("what is the weather"))
|
||||
self.assertEqual(metric.event, "CREATED")
|
||||
|
||||
await module.finish_prompt_metric(metric, 120)
|
||||
|
||||
refreshed = await sync_to_async(PromptMetric.objects.get)(id=metric.id)
|
||||
self.assertEqual(refreshed.event, "FINISHED")
|
||||
self.assertEqual(refreshed.reponse_length, 120)
|
||||
self.assertIsNotNone(refreshed.end_time)
|
||||
|
||||
|
||||
class GraphNodeTestCase(TransactionTestCase):
|
||||
"""The LangGraph nodes behind ``ws/conditional_chat/``."""
|
||||
|
||||
def setUp(self):
|
||||
self.company = make_company()
|
||||
self.user = make_user(company=self.company)
|
||||
self.workspace = make_workspace(self.company)
|
||||
self.conversation = make_conversation(user=self.user)
|
||||
self.prompt = Prompt.objects.create(
|
||||
conversation=self.conversation,
|
||||
message="what is our policy?",
|
||||
user_created=True,
|
||||
)
|
||||
|
||||
def _state(self, **overrides):
|
||||
state = {
|
||||
"message": "what is our policy?",
|
||||
"conversation_id": self.conversation.id,
|
||||
"decoded_file": None,
|
||||
"file_type": None,
|
||||
"messages": [HumanMessage(content="what is our policy?")],
|
||||
"prompt_instance": self.prompt,
|
||||
"moderation_label": ModerationLabel.FINE,
|
||||
"prompt_type": PromptType.GENERAL_CHAT,
|
||||
"response_generator": None,
|
||||
"error": None,
|
||||
"model_name": "Turbo",
|
||||
}
|
||||
state.update(overrides)
|
||||
return state
|
||||
|
||||
async def test_moderation_node_records_label(self):
|
||||
with mock.patch.object(
|
||||
consumers_graph.moderation_classifier,
|
||||
"classify_async",
|
||||
mock.AsyncMock(return_value=ModerationLabel.NSFW),
|
||||
):
|
||||
result = await consumers_graph.moderation_node(self._state())
|
||||
|
||||
self.assertEqual(result["moderation_label"], ModerationLabel.NSFW)
|
||||
|
||||
async def test_classification_node_skipped_for_nsfw_prompts(self):
|
||||
with mock.patch.object(
|
||||
consumers_graph.PROMPT_CLASSIFIER, "classify_async", mock.AsyncMock()
|
||||
) as classify:
|
||||
result = await consumers_graph.classification_node(
|
||||
self._state(moderation_label=ModerationLabel.NSFW)
|
||||
)
|
||||
|
||||
self.assertIsNone(result["prompt_type"])
|
||||
classify.assert_not_called()
|
||||
|
||||
async def test_classification_node_uses_classifier(self):
|
||||
with mock.patch.object(
|
||||
consumers_graph.PROMPT_CLASSIFIER,
|
||||
"classify_async",
|
||||
mock.AsyncMock(return_value=PromptType.RAG),
|
||||
):
|
||||
result = await consumers_graph.classification_node(self._state())
|
||||
|
||||
self.assertEqual(result["prompt_type"], PromptType.RAG)
|
||||
|
||||
async def test_attached_file_forces_data_analysis(self):
|
||||
with mock.patch.object(
|
||||
consumers_graph.PROMPT_CLASSIFIER,
|
||||
"classify_async",
|
||||
mock.AsyncMock(return_value=PromptType.GENERAL_CHAT),
|
||||
):
|
||||
result = await consumers_graph.classification_node(
|
||||
self._state(message="analyze this file", decoded_file=b"a,b\n1,2\n")
|
||||
)
|
||||
|
||||
self.assertEqual(result["prompt_type"], PromptType.DATA_ANALYSIS)
|
||||
|
||||
async def test_attached_file_with_unrelated_prompt_stays_general_chat(self):
|
||||
with mock.patch.object(
|
||||
consumers_graph.PROMPT_CLASSIFIER,
|
||||
"classify_async",
|
||||
mock.AsyncMock(return_value=PromptType.SEARCH),
|
||||
):
|
||||
result = await consumers_graph.classification_node(
|
||||
self._state(message="tell me a joke", decoded_file=b"a,b\n1,2\n")
|
||||
)
|
||||
|
||||
self.assertEqual(result["prompt_type"], PromptType.GENERAL_CHAT)
|
||||
|
||||
async def test_generation_node_refuses_nsfw_prompts(self):
|
||||
result = await consumers_graph.generation_node(
|
||||
self._state(moderation_label=ModerationLabel.NSFW)
|
||||
)
|
||||
|
||||
payload = result["response_generator"]
|
||||
self.assertEqual(payload["type"], "error")
|
||||
self.assertIn("NSFW", payload["content"])
|
||||
|
||||
@override_settings(ALLOW_IMAGE_GENERATION=False)
|
||||
async def test_generation_node_reports_disabled_image_generation(self):
|
||||
result = await consumers_graph.generation_node(
|
||||
self._state(prompt_type=PromptType.IMAGE_GENERATION)
|
||||
)
|
||||
|
||||
self.assertEqual(
|
||||
result["response_generator"],
|
||||
{"type": "text", "content": "Image Generation is disabled."},
|
||||
)
|
||||
|
||||
@override_settings(ALLOW_IMAGE_GENERATION=True)
|
||||
async def test_generation_node_reports_unimplemented_image_generation(self):
|
||||
result = await consumers_graph.generation_node(
|
||||
self._state(prompt_type=PromptType.IMAGE_GENERATION)
|
||||
)
|
||||
|
||||
self.assertIn(
|
||||
"not supported at this time", result["response_generator"]["content"]
|
||||
)
|
||||
|
||||
async def test_generation_node_requires_a_file_for_data_analysis(self):
|
||||
result = await consumers_graph.generation_node(
|
||||
self._state(prompt_type=PromptType.DATA_ANALYSIS)
|
||||
)
|
||||
|
||||
self.assertEqual(
|
||||
result["response_generator"],
|
||||
{
|
||||
"type": "text",
|
||||
"content": "Please upload a file to perform data analysis.",
|
||||
},
|
||||
)
|
||||
|
||||
async def test_generation_node_streams_data_analysis(self):
|
||||
with mock.patch.object(consumers_graph, "AsyncDataAnalysisService") as service:
|
||||
service.return_value.generate_response.return_value = "generator"
|
||||
|
||||
result = await consumers_graph.generation_node(
|
||||
self._state(
|
||||
prompt_type=PromptType.DATA_ANALYSIS,
|
||||
decoded_file=b"a,b\n1,2\n",
|
||||
file_type="csv",
|
||||
)
|
||||
)
|
||||
|
||||
self.assertEqual(result["response_generator"], "generator")
|
||||
service.return_value.generate_response.assert_called_once_with(
|
||||
self.prompt.message, b"a,b\n1,2\n", "csv"
|
||||
)
|
||||
|
||||
async def test_generation_node_routes_rag_prompts_with_workspace(self):
|
||||
with mock.patch.object(consumers_graph, "AsyncRAGService") as service:
|
||||
service.return_value.generate_response.return_value = "generator"
|
||||
|
||||
result = await consumers_graph.generation_node(
|
||||
self._state(prompt_type=PromptType.RAG)
|
||||
)
|
||||
|
||||
self.assertEqual(result["response_generator"], "generator")
|
||||
_args, _kwargs = service.return_value.generate_response.call_args
|
||||
self.assertEqual(_args[2].id, self.workspace.id)
|
||||
|
||||
async def test_generation_node_defaults_to_general_chat(self):
|
||||
with mock.patch.object(consumers_graph, "AsyncLLMService") as service:
|
||||
service.return_value.generate_response.return_value = "generator"
|
||||
|
||||
result = await consumers_graph.generation_node(self._state())
|
||||
|
||||
self.assertEqual(result["response_generator"], "generator")
|
||||
service.return_value.generate_response.assert_called_once()
|
||||
|
||||
@override_settings(ALLOW_INTERNET_ACCESS=True)
|
||||
async def test_search_prompts_append_web_results(self):
|
||||
state = self._state(prompt_type=PromptType.SEARCH)
|
||||
|
||||
with mock.patch.object(consumers_graph, "DuckDuckGoSearchRun") as search:
|
||||
search.return_value.run.return_value = "top result"
|
||||
with mock.patch.object(consumers_graph, "AsyncLLMService"):
|
||||
await consumers_graph.generation_node(state)
|
||||
|
||||
self.assertIn("Search Results: top result", state["messages"][-1].content)
|
||||
|
||||
@override_settings(ALLOW_INTERNET_ACCESS=True)
|
||||
async def test_fast_model_skips_web_search(self):
|
||||
state = self._state(prompt_type=PromptType.SEARCH, model_name="FAST")
|
||||
|
||||
with mock.patch.object(consumers_graph, "DuckDuckGoSearchRun") as search:
|
||||
with mock.patch.object(consumers_graph, "AsyncLLMService"):
|
||||
await consumers_graph.generation_node(state)
|
||||
|
||||
search.assert_not_called()
|
||||
self.assertEqual(len(state["messages"]), 1)
|
||||
|
||||
@override_settings(ALLOW_INTERNET_ACCESS=False)
|
||||
async def test_search_is_skipped_when_internet_access_is_disabled(self):
|
||||
state = self._state(prompt_type=PromptType.SEARCH)
|
||||
|
||||
with mock.patch.object(consumers_graph, "DuckDuckGoSearchRun") as search:
|
||||
with mock.patch.object(consumers_graph, "AsyncLLMService"):
|
||||
await consumers_graph.generation_node(state)
|
||||
|
||||
search.assert_not_called()
|
||||
|
||||
@override_settings(ALLOW_INTERNET_ACCESS=True)
|
||||
async def test_search_failures_fall_back_to_plain_chat(self):
|
||||
state = self._state(prompt_type=PromptType.SEARCH)
|
||||
|
||||
with mock.patch.object(consumers_graph, "DuckDuckGoSearchRun") as search:
|
||||
search.return_value.run.side_effect = RuntimeError("ddg unreachable")
|
||||
with mock.patch.object(consumers_graph, "AsyncLLMService") as service:
|
||||
service.return_value.generate_response.return_value = "generator"
|
||||
result = await consumers_graph.generation_node(state)
|
||||
|
||||
self.assertEqual(result["response_generator"], "generator")
|
||||
self.assertEqual(len(state["messages"]), 1)
|
||||
|
||||
|
||||
class WebSocketRoutingTestCase(TransactionTestCase):
|
||||
@parameterized.expand(
|
||||
[("chat", "/ws/chat_again/"), ("conditional_chat", "/ws/conditional_chat/")]
|
||||
)
|
||||
async def test_route_accepts_connections(self, _name, path):
|
||||
communicator = WebsocketCommunicator(application, path)
|
||||
|
||||
connected, _subprotocol = await communicator.connect()
|
||||
self.assertTrue(connected)
|
||||
|
||||
await communicator.disconnect()
|
||||
|
||||
def test_only_the_two_chat_routes_are_registered(self):
|
||||
self.assertEqual(
|
||||
[str(route.pattern) for route in websocket_urlpatterns],
|
||||
["ws/chat_again/$", "ws/conditional_chat/$"],
|
||||
)
|
||||
@@ -0,0 +1,72 @@
|
||||
"""Opt-in checks that talk to a real Ollama instance.
|
||||
|
||||
These are excluded from CI because they need a model server and their answers are
|
||||
not deterministic. Run them against a reachable ``OLLAMA_BASE_URL`` with::
|
||||
|
||||
RUN_LIVE_OLLAMA_TESTS=1 uv run python manage.py test \
|
||||
chat_backend.tests.test_live_ollama
|
||||
"""
|
||||
|
||||
import os
|
||||
import unittest
|
||||
|
||||
from parameterized import parameterized
|
||||
|
||||
from chat_backend.services.moderation_classifier import (
|
||||
ModerationClassifier,
|
||||
ModerationLabel,
|
||||
)
|
||||
from chat_backend.services.prompt_classifier.prompt_classifier import (
|
||||
PromptClassifier,
|
||||
PromptType,
|
||||
)
|
||||
from chat_backend.services.title_generator import TitleGenerator
|
||||
|
||||
live_only = unittest.skipUnless(
|
||||
os.environ.get("RUN_LIVE_OLLAMA_TESTS", "").lower() in {"1", "true", "yes"},
|
||||
"Requires a live Ollama; set RUN_LIVE_OLLAMA_TESTS=1 to run",
|
||||
)
|
||||
|
||||
|
||||
@live_only
|
||||
class PromptClassifierLiveTestCase(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.service = PromptClassifier()
|
||||
|
||||
@parameterized.expand(
|
||||
[
|
||||
["Tell me a joke", PromptType.GENERAL_CHAT],
|
||||
["Create an image of a dog for me", PromptType.IMAGE_GENERATION],
|
||||
[
|
||||
"highlight the features of the backyard playset if they were to choose "
|
||||
"us and make the language more long form",
|
||||
PromptType.GENERAL_CHAT,
|
||||
],
|
||||
]
|
||||
)
|
||||
def test_prompt_classification(self, prompt, expected_output):
|
||||
self.assertEqual(self.service.classify(prompt), expected_output)
|
||||
|
||||
|
||||
@live_only
|
||||
class ModerationClassifierLiveTestCase(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.service = ModerationClassifier()
|
||||
|
||||
@parameterized.expand(
|
||||
[
|
||||
["Write a love poem", ModerationLabel.FINE],
|
||||
["How to make a bomb", ModerationLabel.NSFW],
|
||||
]
|
||||
)
|
||||
def test_moderation(self, prompt, expected_output):
|
||||
self.assertEqual(self.service.classify(prompt), expected_output)
|
||||
|
||||
|
||||
@live_only
|
||||
class TitleGeneratorLiveTestCase(unittest.TestCase):
|
||||
def test_title_is_short(self):
|
||||
title = TitleGenerator().generate("What's the weather today?")
|
||||
|
||||
self.assertTrue(title)
|
||||
self.assertLessEqual(len(title.split()), 5)
|
||||
@@ -0,0 +1,194 @@
|
||||
import datetime
|
||||
|
||||
from django.core.files.uploadedfile import SimpleUploadedFile
|
||||
from django.test import TestCase
|
||||
from django.utils import timezone
|
||||
|
||||
from chat_backend.models import (
|
||||
Announcement,
|
||||
Document,
|
||||
Feedback,
|
||||
LLMModels,
|
||||
Prompt,
|
||||
PromptMetric,
|
||||
StoredFile,
|
||||
)
|
||||
|
||||
from .factories import (
|
||||
make_company,
|
||||
make_conversation,
|
||||
make_document,
|
||||
make_prompt,
|
||||
make_user,
|
||||
make_workspace,
|
||||
)
|
||||
|
||||
|
||||
class TimeInfoBaseTestCase(TestCase):
|
||||
def test_save_refreshes_last_modified(self):
|
||||
company = make_company()
|
||||
original = company.last_modified
|
||||
|
||||
company.name = "Renamed"
|
||||
company.save()
|
||||
|
||||
company.refresh_from_db()
|
||||
self.assertGreater(company.last_modified, original)
|
||||
self.assertEqual(company.name, "Renamed")
|
||||
|
||||
def test_skip_last_modified_keeps_timestamp(self):
|
||||
company = make_company()
|
||||
original = company.last_modified
|
||||
|
||||
company.name = "Renamed"
|
||||
company.save(skip_last_modified=True)
|
||||
|
||||
company.refresh_from_db()
|
||||
self.assertEqual(company.last_modified, original)
|
||||
self.assertEqual(company.name, "Renamed")
|
||||
|
||||
def test_update_fields_includes_last_modified(self):
|
||||
conversation = make_conversation(title="Before")
|
||||
original = conversation.last_modified
|
||||
|
||||
conversation.title = "After"
|
||||
conversation.save(update_fields=["title"])
|
||||
|
||||
conversation.refresh_from_db()
|
||||
self.assertEqual(conversation.title, "After")
|
||||
self.assertGreater(conversation.last_modified, original)
|
||||
|
||||
|
||||
class CompanyAndUserTestCase(TestCase):
|
||||
def test_company_str_is_name(self):
|
||||
self.assertEqual(str(make_company("Globex")), "Globex")
|
||||
|
||||
def test_user_slug_is_generated_from_email(self):
|
||||
user = make_user(email="person@example.com")
|
||||
self.assertTrue(user.slug)
|
||||
self.assertNotIn("@", user.slug)
|
||||
|
||||
def test_set_password_url_contains_slug(self):
|
||||
user = make_user(email="person@example.com")
|
||||
self.assertEqual(
|
||||
user.get_set_password_url(),
|
||||
f"https://chat.aimloperations.com/set_password?slug={user.slug}",
|
||||
)
|
||||
|
||||
def test_user_defaults(self):
|
||||
user = make_user()
|
||||
self.assertFalse(user.is_company_manager)
|
||||
self.assertFalse(user.deleted)
|
||||
self.assertFalse(user.has_signed_tos)
|
||||
self.assertTrue(user.conversation_order)
|
||||
|
||||
|
||||
class ConversationAndPromptTestCase(TestCase):
|
||||
def setUp(self):
|
||||
self.company = make_company()
|
||||
self.user = make_user(company=self.company)
|
||||
self.conversation = make_conversation(user=self.user, title="Weather Inquiry")
|
||||
|
||||
def test_conversation_str_and_user_email(self):
|
||||
self.assertEqual(str(self.conversation), "Weather Inquiry")
|
||||
self.assertEqual(self.conversation.get_user_email(), self.user.email)
|
||||
|
||||
def test_conversation_without_user_returns_empty_email(self):
|
||||
self.assertEqual(make_conversation().get_user_email(), "")
|
||||
|
||||
def test_prompt_conversation_title(self):
|
||||
prompt = make_prompt(self.conversation)
|
||||
self.assertEqual(prompt.get_conversation_title(), "Weather Inquiry")
|
||||
|
||||
def test_prompt_without_conversation_returns_empty_title(self):
|
||||
prompt = Prompt.objects.create(message="orphan", user_created=True)
|
||||
self.assertEqual(prompt.get_conversation_title(), "")
|
||||
|
||||
def test_file_exists_false_without_file(self):
|
||||
self.assertFalse(make_prompt(self.conversation).file_exists())
|
||||
|
||||
def test_file_exists_tracks_stored_blob(self):
|
||||
prompt = make_prompt(self.conversation)
|
||||
prompt.file.save("data.csv", SimpleUploadedFile("data.csv", b"a,b\n1,2\n"))
|
||||
|
||||
self.assertTrue(prompt.file_exists())
|
||||
self.assertTrue(StoredFile.objects.filter(name=prompt.file.name).exists())
|
||||
|
||||
StoredFile.objects.filter(name=prompt.file.name).delete()
|
||||
self.assertFalse(prompt.file_exists())
|
||||
|
||||
|
||||
class PromptMetricTestCase(TestCase):
|
||||
def _metric(self, **kwargs):
|
||||
defaults = {
|
||||
"prompt_id": 1,
|
||||
"conversation_id": 1,
|
||||
"model_name": "llama3.2",
|
||||
"start_time": timezone.now(),
|
||||
"prompt_length": 10,
|
||||
"has_file": False,
|
||||
}
|
||||
defaults.update(kwargs)
|
||||
return PromptMetric.objects.create(**defaults)
|
||||
|
||||
def test_duration_is_zero_until_finished(self):
|
||||
self.assertEqual(self._metric().get_duration(), 0)
|
||||
|
||||
def test_duration_in_seconds(self):
|
||||
start = timezone.now()
|
||||
metric = self._metric(
|
||||
start_time=start, end_time=start + datetime.timedelta(seconds=42)
|
||||
)
|
||||
self.assertEqual(metric.get_duration(), 42)
|
||||
|
||||
def test_default_event_is_created(self):
|
||||
self.assertEqual(self._metric().event, "CREATED")
|
||||
|
||||
|
||||
class FeedbackTestCase(TestCase):
|
||||
def test_get_user_email(self):
|
||||
user = make_user()
|
||||
feedback = Feedback.objects.create(title="t", text="body", user=user)
|
||||
self.assertEqual(feedback.get_user_email(), user.email)
|
||||
|
||||
def test_get_user_email_without_user(self):
|
||||
feedback = Feedback.objects.create(title="t", text="body")
|
||||
self.assertEqual(feedback.get_user_email(), "")
|
||||
|
||||
def test_defaults(self):
|
||||
feedback = Feedback.objects.create(text="body")
|
||||
self.assertEqual(feedback.status, "SUBMITTED")
|
||||
self.assertEqual(feedback.category, "NOT_DEFINED")
|
||||
|
||||
|
||||
class DocumentModelTestCase(TestCase):
|
||||
def setUp(self):
|
||||
self.workspace = make_workspace(make_company())
|
||||
|
||||
def test_document_defaults_and_storage(self):
|
||||
document = make_document(self.workspace)
|
||||
|
||||
self.assertFalse(document.processed)
|
||||
self.assertFalse(document.active)
|
||||
self.assertIsNotNone(document.uploaded_at)
|
||||
self.assertTrue(document.file.name.startswith("documents/"))
|
||||
self.assertTrue(StoredFile.objects.filter(name=document.file.name).exists())
|
||||
|
||||
def test_deleting_workspace_cascades_to_documents(self):
|
||||
make_document(self.workspace)
|
||||
self.workspace.delete()
|
||||
self.assertEqual(Document.objects.count(), 0)
|
||||
|
||||
|
||||
class MiscModelTestCase(TestCase):
|
||||
def test_stored_file_str_is_name(self):
|
||||
stored = StoredFile.objects.create(name="documents/a.txt", content=b"x", size=1)
|
||||
self.assertEqual(str(stored), "documents/a.txt")
|
||||
|
||||
def test_announcement_default_status(self):
|
||||
announcement = Announcement.objects.create(message="hello")
|
||||
self.assertEqual(announcement.status, Announcement.Status.default)
|
||||
|
||||
def test_llm_model_fields(self):
|
||||
model = LLMModels.objects.create(name="llama3.2", port=11434, description="d")
|
||||
self.assertEqual(model.port, 11434)
|
||||
@@ -0,0 +1,188 @@
|
||||
from django.test import TestCase
|
||||
|
||||
from chat_backend.models import Prompt
|
||||
from chat_backend.serializers import (
|
||||
BasicUserSerializer,
|
||||
ConversationSerializer,
|
||||
CustomUserSerializer,
|
||||
DocumentSerializer,
|
||||
DocumentWorkspaceSerializer,
|
||||
FeedbackSerializer,
|
||||
MyTokenObtainPairSerializer,
|
||||
PromptSerializer,
|
||||
)
|
||||
|
||||
from .factories import (
|
||||
make_company,
|
||||
make_conversation,
|
||||
make_document,
|
||||
make_user,
|
||||
make_workspace,
|
||||
pdf_upload,
|
||||
)
|
||||
|
||||
|
||||
class TokenSerializerTestCase(TestCase):
|
||||
def test_token_carries_company_claim(self):
|
||||
user = make_user(company=make_company())
|
||||
|
||||
token = MyTokenObtainPairSerializer.get_token(user)
|
||||
|
||||
self.assertEqual(token["company"], "something here")
|
||||
self.assertEqual(str(token["user_id"]), str(user.id))
|
||||
|
||||
|
||||
class ConversationSerializerTestCase(TestCase):
|
||||
def test_exposes_only_whitelisted_fields(self):
|
||||
conversation = make_conversation(user=make_user(), title="Weather Inquiry")
|
||||
|
||||
data = ConversationSerializer(conversation).data
|
||||
|
||||
self.assertEqual(set(data.keys()), {"title", "created", "last_modified", "id"})
|
||||
self.assertEqual(data["title"], "Weather Inquiry")
|
||||
|
||||
|
||||
class PromptSerializerTestCase(TestCase):
|
||||
def test_valid_payload_creates_prompt(self):
|
||||
serializer = PromptSerializer(data={"message": "hi", "user_created": True})
|
||||
|
||||
self.assertTrue(serializer.is_valid(), serializer.errors)
|
||||
prompt = serializer.save()
|
||||
|
||||
self.assertEqual(Prompt.objects.count(), 1)
|
||||
self.assertTrue(prompt.user_created)
|
||||
|
||||
def test_message_is_required(self):
|
||||
serializer = PromptSerializer(data={"user_created": True})
|
||||
|
||||
self.assertFalse(serializer.is_valid())
|
||||
self.assertIn("message", serializer.errors)
|
||||
|
||||
def test_user_created_is_required(self):
|
||||
serializer = PromptSerializer(data={"message": "hi"})
|
||||
|
||||
self.assertFalse(serializer.is_valid())
|
||||
self.assertIn("user_created", serializer.errors)
|
||||
|
||||
def test_file_is_not_serialized(self):
|
||||
prompt = Prompt.objects.create(message="hi", user_created=True)
|
||||
self.assertNotIn("file", PromptSerializer(prompt).data)
|
||||
|
||||
|
||||
class UserSerializerTestCase(TestCase):
|
||||
def setUp(self):
|
||||
self.user = make_user(company=make_company("Globex"), first_name="Ada")
|
||||
|
||||
def test_custom_user_serializer_hides_password(self):
|
||||
data = CustomUserSerializer(self.user).data
|
||||
|
||||
self.assertEqual(data["email"], self.user.email)
|
||||
self.assertEqual(data["company"]["name"], "Globex")
|
||||
self.assertTrue(data["has_usable_password"])
|
||||
self.assertNotIn("password", data)
|
||||
|
||||
def test_custom_user_serializer_rejects_short_password(self):
|
||||
serializer = CustomUserSerializer(
|
||||
data={
|
||||
"username": "someone",
|
||||
"email": "someone@example.com",
|
||||
"password": "short",
|
||||
"has_usable_password": True,
|
||||
"company": {
|
||||
"name": "Acme",
|
||||
"state": "IL",
|
||||
"zipcode": "60189",
|
||||
"address": "1 Main St",
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
self.assertFalse(serializer.is_valid())
|
||||
self.assertIn("password", serializer.errors)
|
||||
|
||||
def test_custom_user_serializer_rejects_invalid_email(self):
|
||||
serializer = CustomUserSerializer(
|
||||
data={
|
||||
"username": "someone",
|
||||
"email": "not-an-email",
|
||||
"password": "longenough",
|
||||
}
|
||||
)
|
||||
|
||||
self.assertFalse(serializer.is_valid())
|
||||
self.assertIn("email", serializer.errors)
|
||||
|
||||
def test_basic_user_serializer_field_set(self):
|
||||
data = BasicUserSerializer(self.user).data
|
||||
|
||||
self.assertEqual(
|
||||
set(data.keys()),
|
||||
{
|
||||
"email",
|
||||
"first_name",
|
||||
"last_name",
|
||||
"is_active",
|
||||
"has_usable_password",
|
||||
"is_company_manager",
|
||||
"has_signed_tos",
|
||||
},
|
||||
)
|
||||
self.assertEqual(data["first_name"], "Ada")
|
||||
|
||||
|
||||
class FeedbackSerializerTestCase(TestCase):
|
||||
def test_text_is_required(self):
|
||||
serializer = FeedbackSerializer(data={"title": "no body"})
|
||||
|
||||
self.assertFalse(serializer.is_valid())
|
||||
self.assertIn("text", serializer.errors)
|
||||
|
||||
def test_defaults_applied_on_save(self):
|
||||
serializer = FeedbackSerializer(data={"title": "t", "text": "body"})
|
||||
|
||||
self.assertTrue(serializer.is_valid(), serializer.errors)
|
||||
feedback = serializer.save()
|
||||
|
||||
self.assertEqual(feedback.status, "SUBMITTED")
|
||||
self.assertEqual(feedback.category, "NOT_DEFINED")
|
||||
|
||||
|
||||
class DocumentSerializerTestCase(TestCase):
|
||||
def setUp(self):
|
||||
self.workspace = make_workspace(make_company())
|
||||
|
||||
def test_workspace_serializer_read_only_fields(self):
|
||||
serializer = DocumentWorkspaceSerializer(
|
||||
data={"id": 999, "name": "New Workspace"}
|
||||
)
|
||||
|
||||
self.assertTrue(serializer.is_valid(), serializer.errors)
|
||||
self.assertEqual(set(serializer.validated_data.keys()), {"name"})
|
||||
|
||||
def test_workspace_name_is_required(self):
|
||||
serializer = DocumentWorkspaceSerializer(data={})
|
||||
|
||||
self.assertFalse(serializer.is_valid())
|
||||
self.assertIn("name", serializer.errors)
|
||||
|
||||
def test_document_serializer_reports_stored_file_url(self):
|
||||
document = make_document(self.workspace)
|
||||
|
||||
data = DocumentSerializer(document).data
|
||||
|
||||
self.assertEqual(data["workspace"], self.workspace.id)
|
||||
self.assertIn("test", data["file"])
|
||||
self.assertIn("pdf", data["file"])
|
||||
self.assertFalse(data["processed"])
|
||||
|
||||
def test_document_serializer_ignores_read_only_processed(self):
|
||||
serializer = DocumentSerializer(
|
||||
data={
|
||||
"workspace": self.workspace.id,
|
||||
"file": pdf_upload(),
|
||||
"processed": True,
|
||||
}
|
||||
)
|
||||
|
||||
self.assertTrue(serializer.is_valid(), serializer.errors)
|
||||
self.assertNotIn("processed", serializer.validated_data)
|
||||
@@ -0,0 +1,221 @@
|
||||
from django.test import SimpleTestCase, override_settings
|
||||
from parameterized import parameterized
|
||||
|
||||
from chat_backend.services.moderation_classifier import (
|
||||
ModerationClassifier,
|
||||
ModerationLabel,
|
||||
)
|
||||
from chat_backend.services.prompt_classifier.prompt_classifier import (
|
||||
PromptClassifier,
|
||||
PromptType,
|
||||
)
|
||||
from chat_backend.services.title_generator import TitleGenerator
|
||||
|
||||
from .fakes import FakeChain
|
||||
|
||||
|
||||
class PromptClassifierQuickCheckTestCase(SimpleTestCase):
|
||||
def setUp(self):
|
||||
self.classifier = PromptClassifier()
|
||||
|
||||
@parameterized.expand(
|
||||
[
|
||||
(
|
||||
"data_analysis",
|
||||
"Please read this document and make me an index",
|
||||
PromptType.DATA_ANALYSIS,
|
||||
),
|
||||
(
|
||||
"csv",
|
||||
"Based on this csv, who bought the most?",
|
||||
PromptType.DATA_ANALYSIS,
|
||||
),
|
||||
("rag", "What does the uploaded pdf say about revenue?", PromptType.RAG),
|
||||
(
|
||||
"search_news",
|
||||
"What is the latest news on the merger?",
|
||||
PromptType.SEARCH,
|
||||
),
|
||||
("search_sports", "Who won the game last night?", PromptType.SEARCH),
|
||||
]
|
||||
)
|
||||
def test_rule_based_hits(self, _name, prompt, expected):
|
||||
self.assertEqual(self.classifier._quick_check(prompt), expected)
|
||||
|
||||
def test_plain_chat_falls_through_to_the_llm(self):
|
||||
self.assertIsNone(self.classifier._quick_check("Tell me a joke"))
|
||||
|
||||
@override_settings(ALLOW_IMAGE_GENERATION=True)
|
||||
def test_image_requests_when_feature_enabled(self):
|
||||
self.assertEqual(
|
||||
self.classifier._quick_check("Generate an image of a dog"),
|
||||
PromptType.IMAGE_GENERATION,
|
||||
)
|
||||
|
||||
@override_settings(ALLOW_IMAGE_GENERATION=False)
|
||||
def test_image_requests_downgrade_when_feature_disabled(self):
|
||||
self.assertEqual(
|
||||
self.classifier._quick_check("Generate an image of a dog"),
|
||||
PromptType.GENERAL_CHAT,
|
||||
)
|
||||
|
||||
|
||||
class PromptClassifierParseResponseTestCase(SimpleTestCase):
|
||||
def setUp(self):
|
||||
self.classifier = PromptClassifier()
|
||||
|
||||
@parameterized.expand(
|
||||
[
|
||||
("exact", "RAG", PromptType.RAG),
|
||||
("lowercase", "general_chat", PromptType.GENERAL_CHAT),
|
||||
("missing_underscore", "GENERALCHAT", PromptType.GENERAL_CHAT),
|
||||
("padded", " SEARCH ", PromptType.SEARCH),
|
||||
("verbose", "The answer is DATA_ANALYSIS", PromptType.DATA_ANALYSIS),
|
||||
("nonsense", "banana", PromptType.UNKNOWN),
|
||||
]
|
||||
)
|
||||
def test_parses_llm_labels(self, _name, response, expected):
|
||||
self.assertEqual(self.classifier._parse_response(response), expected)
|
||||
|
||||
@override_settings(ALLOW_IMAGE_GENERATION=False)
|
||||
def test_image_generation_is_downgraded_when_disabled(self):
|
||||
self.assertEqual(
|
||||
self.classifier._parse_response("IMAGE_GENERATION"), PromptType.GENERAL_CHAT
|
||||
)
|
||||
|
||||
@override_settings(ALLOW_IMAGE_GENERATION=True)
|
||||
def test_image_generation_survives_when_enabled(self):
|
||||
self.assertEqual(
|
||||
self.classifier._parse_response("IMAGE_GENERATION"),
|
||||
PromptType.IMAGE_GENERATION,
|
||||
)
|
||||
|
||||
|
||||
class PromptClassifierClassifyTestCase(SimpleTestCase):
|
||||
def setUp(self):
|
||||
self.classifier = PromptClassifier()
|
||||
|
||||
def test_classify_uses_llm_when_no_rule_matches(self):
|
||||
self.classifier.chain = FakeChain("GENERAL_CHAT")
|
||||
|
||||
result = self.classifier.classify("Tell me a joke")
|
||||
|
||||
self.assertEqual(result, PromptType.GENERAL_CHAT)
|
||||
self.assertEqual(self.classifier.chain.calls, [{"prompt": "Tell me a joke"}])
|
||||
|
||||
def test_classify_skips_llm_when_rule_matches(self):
|
||||
self.classifier.chain = FakeChain("GENERAL_CHAT")
|
||||
|
||||
result = self.classifier.classify("What is the latest news?")
|
||||
|
||||
self.assertEqual(result, PromptType.SEARCH)
|
||||
self.assertEqual(self.classifier.chain.calls, [])
|
||||
|
||||
def test_classify_returns_unknown_when_llm_fails(self):
|
||||
self.classifier.chain = FakeChain(error=RuntimeError("ollama down"))
|
||||
|
||||
self.assertEqual(self.classifier.classify("Tell me a joke"), PromptType.UNKNOWN)
|
||||
|
||||
async def test_classify_async_uses_llm(self):
|
||||
self.classifier.chain = FakeChain("RAG")
|
||||
|
||||
self.assertEqual(
|
||||
await self.classifier.classify_async("Tell me a joke"), PromptType.RAG
|
||||
)
|
||||
|
||||
async def test_classify_async_returns_unknown_when_llm_fails(self):
|
||||
self.classifier.chain = FakeChain(error=RuntimeError("ollama down"))
|
||||
|
||||
self.assertEqual(
|
||||
await self.classifier.classify_async("Tell me a joke"), PromptType.UNKNOWN
|
||||
)
|
||||
|
||||
|
||||
class ModerationClassifierTestCase(SimpleTestCase):
|
||||
def setUp(self):
|
||||
self.classifier = ModerationClassifier()
|
||||
|
||||
@parameterized.expand(
|
||||
[
|
||||
("nsfw", "NSFW", ModerationLabel.NSFW),
|
||||
("fine", "FINE", ModerationLabel.FINE),
|
||||
("verbose_nsfw", "This prompt is NSFW because...", ModerationLabel.NSFW),
|
||||
("unclear_defaults_to_fine", "not sure", ModerationLabel.FINE),
|
||||
]
|
||||
)
|
||||
def test_parses_labels(self, _name, response, expected):
|
||||
self.assertEqual(self.classifier._parse_response(response), expected)
|
||||
|
||||
def test_classify_normalizes_llm_output(self):
|
||||
self.classifier.chain = FakeChain(" fine \n")
|
||||
|
||||
self.assertEqual(
|
||||
self.classifier.classify("Python tutorial"), ModerationLabel.FINE
|
||||
)
|
||||
|
||||
def test_classify_fails_safe_to_nsfw(self):
|
||||
self.classifier.chain = FakeChain(error=RuntimeError("ollama down"))
|
||||
|
||||
self.assertEqual(
|
||||
self.classifier.classify("Python tutorial"), ModerationLabel.NSFW
|
||||
)
|
||||
|
||||
async def test_classify_async_fails_safe_to_nsfw(self):
|
||||
self.classifier.chain = FakeChain(error=RuntimeError("ollama down"))
|
||||
|
||||
self.assertEqual(
|
||||
await self.classifier.classify_async("Python tutorial"),
|
||||
ModerationLabel.NSFW,
|
||||
)
|
||||
|
||||
|
||||
class TitleGeneratorTestCase(SimpleTestCase):
|
||||
def setUp(self):
|
||||
self.generator = TitleGenerator()
|
||||
|
||||
@parameterized.expand(
|
||||
[
|
||||
("strips_quotes", '"Weather Inquiry"', "Weather Inquiry"),
|
||||
("strips_punctuation", "weather inquiry!", "Weather Inquiry"),
|
||||
(
|
||||
"title_cases",
|
||||
"quantum computing explanation",
|
||||
"Quantum Computing Explanation",
|
||||
),
|
||||
]
|
||||
)
|
||||
def test_cleans_llm_output(self, _name, response, expected):
|
||||
self.assertEqual(self.generator._clean_response(response), expected)
|
||||
|
||||
def test_titles_are_capped_at_fifty_characters(self):
|
||||
cleaned = self.generator._clean_response("word " * 40)
|
||||
|
||||
self.assertEqual(len(cleaned), 50)
|
||||
|
||||
def test_generate_returns_cleaned_title(self):
|
||||
self.generator.chain = FakeChain('"dragon image generation"')
|
||||
|
||||
self.assertEqual(
|
||||
self.generator.generate("Generate an image of a dragon"),
|
||||
"Dragon Image Generation",
|
||||
)
|
||||
|
||||
def test_generate_falls_back_on_error(self):
|
||||
self.generator.chain = FakeChain(error=RuntimeError("ollama down"))
|
||||
|
||||
self.assertEqual(self.generator.generate("anything"), "Conversation")
|
||||
|
||||
async def test_generate_async_falls_back_on_error(self):
|
||||
self.generator.chain = FakeChain(error=RuntimeError("ollama down"))
|
||||
|
||||
self.assertEqual(
|
||||
await self.generator.generate_async("anything"), "Conversation"
|
||||
)
|
||||
|
||||
async def test_generate_async_returns_cleaned_title(self):
|
||||
self.generator.chain = FakeChain("weather inquiry")
|
||||
|
||||
self.assertEqual(
|
||||
await self.generator.generate_async("What's the weather?"),
|
||||
"Weather Inquiry",
|
||||
)
|
||||
@@ -0,0 +1,173 @@
|
||||
import base64
|
||||
import json
|
||||
|
||||
import pandas as pd
|
||||
from django.test import SimpleTestCase
|
||||
|
||||
from chat_backend.services.data_analysis_service import AsyncDataAnalysisService
|
||||
|
||||
from .factories import CSV_BYTES, docx_bytes_with_text, pdf_bytes_with_text
|
||||
from .fakes import FakeChain
|
||||
|
||||
PNG_MAGIC = b"\x89PNG\r\n\x1a\n"
|
||||
|
||||
|
||||
def frame():
|
||||
return pd.DataFrame(
|
||||
{"name": ["a", "b", "c"], "sales": [1, 2, 3], "units": [4, 5, 6]}
|
||||
)
|
||||
|
||||
|
||||
async def collect(generator):
|
||||
return [chunk async for chunk in generator]
|
||||
|
||||
|
||||
class DataFrameSummaryTestCase(SimpleTestCase):
|
||||
def setUp(self):
|
||||
self.service = AsyncDataAnalysisService()
|
||||
|
||||
def test_summary_reports_shape_columns_and_sample(self):
|
||||
summary = self.service._get_dataframe_summary(frame())
|
||||
|
||||
self.assertIn("DataFrame has 3 rows and 3 columns.", summary)
|
||||
self.assertIn("Descriptive Statistics", summary)
|
||||
self.assertIn("sales", summary)
|
||||
self.assertIn("units", summary)
|
||||
|
||||
|
||||
class DocumentReadersTestCase(SimpleTestCase):
|
||||
def setUp(self):
|
||||
self.service = AsyncDataAnalysisService()
|
||||
|
||||
def test_reads_docx_paragraphs(self):
|
||||
text = self.service._read_docx(
|
||||
docx_bytes_with_text("first line", "second line")
|
||||
)
|
||||
|
||||
self.assertEqual(text, "first line\nsecond line")
|
||||
|
||||
def test_reads_pdf_text(self):
|
||||
text = self.service._read_pdf(pdf_bytes_with_text("Quarterly Report"))
|
||||
|
||||
self.assertIn("Quarterly Report", text)
|
||||
|
||||
|
||||
class PlotGenerationTestCase(SimpleTestCase):
|
||||
def setUp(self):
|
||||
self.service = AsyncDataAnalysisService()
|
||||
|
||||
def test_uses_columns_named_in_the_query(self):
|
||||
encoded = self.service._generate_plot("plot sales vs units", frame())
|
||||
|
||||
self.assertTrue(base64.b64decode(encoded).startswith(PNG_MAGIC))
|
||||
|
||||
def test_falls_back_to_first_two_numeric_columns(self):
|
||||
encoded = self.service._generate_plot("graph the data please", frame())
|
||||
|
||||
self.assertTrue(base64.b64decode(encoded).startswith(PNG_MAGIC))
|
||||
|
||||
def test_requires_two_numeric_columns(self):
|
||||
single_numeric = pd.DataFrame({"name": ["a", "b"], "sales": [1, 2]})
|
||||
|
||||
with self.assertRaises(ValueError):
|
||||
self.service._generate_plot("plot it", single_numeric)
|
||||
|
||||
|
||||
class GenerateResponseTestCase(SimpleTestCase):
|
||||
def setUp(self):
|
||||
self.service = AsyncDataAnalysisService()
|
||||
self.service.analysis_chain = FakeChain(chunks=["Sales ", "are up."])
|
||||
|
||||
async def test_csv_question_is_answered_from_the_summary(self):
|
||||
chunks = await collect(
|
||||
self.service.generate_response(
|
||||
"What are total sales?", CSV_BYTES, "text/csv"
|
||||
)
|
||||
)
|
||||
|
||||
self.assertEqual("".join(chunks), "Sales are up.")
|
||||
payload = self.service.analysis_chain.calls[0]
|
||||
self.assertEqual(payload["query"], "What are total sales?")
|
||||
self.assertIn("DataFrame has 3 rows", payload["data_summary"])
|
||||
|
||||
async def test_plot_request_returns_png_payload(self):
|
||||
chunks = await collect(
|
||||
self.service.generate_response("plot sales vs units", CSV_BYTES, "text/csv")
|
||||
)
|
||||
|
||||
self.assertEqual(len(chunks), 1)
|
||||
payload = json.loads(chunks[0])
|
||||
self.assertEqual(payload["type"], "plot")
|
||||
self.assertEqual(payload["format"], "png")
|
||||
self.assertTrue(base64.b64decode(payload["image"]).startswith(PNG_MAGIC))
|
||||
self.assertEqual(self.service.analysis_chain.calls, [])
|
||||
|
||||
async def test_plot_request_without_numeric_columns_reports_error(self):
|
||||
csv = b"name,city\nalice,chicago\nbob,denver\n"
|
||||
|
||||
chunks = await collect(
|
||||
self.service.generate_response("plot it", csv, "text/csv")
|
||||
)
|
||||
|
||||
payload = json.loads(chunks[0])
|
||||
self.assertEqual(payload["type"], "error")
|
||||
self.assertIn("two numerical columns", payload["content"])
|
||||
|
||||
async def test_docx_question_uses_document_text(self):
|
||||
docx = docx_bytes_with_text("Revenue grew by 10%")
|
||||
|
||||
await collect(
|
||||
self.service.generate_response(
|
||||
"Summarize this",
|
||||
docx,
|
||||
"application/vnd.openxmlformats-officedocument.wordprocessingml.document",
|
||||
)
|
||||
)
|
||||
|
||||
self.assertIn(
|
||||
"Revenue grew by 10%", self.service.analysis_chain.calls[0]["data_summary"]
|
||||
)
|
||||
|
||||
async def test_pdf_question_uses_document_text(self):
|
||||
await collect(
|
||||
self.service.generate_response(
|
||||
"Summarize this",
|
||||
pdf_bytes_with_text("Quarterly Report"),
|
||||
"application/pdf",
|
||||
)
|
||||
)
|
||||
|
||||
self.assertIn(
|
||||
"Quarterly Report", self.service.analysis_chain.calls[0]["data_summary"]
|
||||
)
|
||||
|
||||
async def test_unsupported_file_type_is_reported(self):
|
||||
chunks = await collect(
|
||||
self.service.generate_response("What is this?", b"\x00\x01", "image/png")
|
||||
)
|
||||
|
||||
payload = json.loads(chunks[0])
|
||||
self.assertEqual(payload["type"], "error")
|
||||
self.assertIn("Unsupported file type", payload["content"])
|
||||
|
||||
async def test_unreadable_file_is_reported(self):
|
||||
chunks = await collect(
|
||||
self.service.generate_response("What is this?", b"", "text/csv")
|
||||
)
|
||||
|
||||
payload = json.loads(chunks[0])
|
||||
self.assertEqual(payload["type"], "error")
|
||||
self.assertIn("An error occurred", payload["content"])
|
||||
|
||||
async def test_llm_failure_is_reported(self):
|
||||
self.service.analysis_chain = FakeChain(error=RuntimeError("ollama down"))
|
||||
|
||||
chunks = await collect(
|
||||
self.service.generate_response(
|
||||
"What are total sales?", CSV_BYTES, "text/csv"
|
||||
)
|
||||
)
|
||||
|
||||
payload = json.loads(chunks[0])
|
||||
self.assertEqual(payload["type"], "error")
|
||||
self.assertIn("ollama down", payload["content"])
|
||||
@@ -0,0 +1,64 @@
|
||||
from django.test import SimpleTestCase
|
||||
from langchain_core.messages import AIMessage, HumanMessage
|
||||
|
||||
from chat_backend.services.llm_service import AsyncLLMService, SyncLLMService
|
||||
|
||||
from .fakes import FakeChain
|
||||
|
||||
|
||||
def conversation(length: int):
|
||||
messages = []
|
||||
for index in range(length):
|
||||
messages.append(HumanMessage(content=f"question {index}"))
|
||||
messages.append(AIMessage(content=f"answer {index}"))
|
||||
return messages
|
||||
|
||||
|
||||
class AsyncLLMServiceTestCase(SimpleTestCase):
|
||||
def setUp(self):
|
||||
self.service = AsyncLLMService()
|
||||
|
||||
async def test_format_history_labels_speakers(self):
|
||||
history = await self.service._format_history(
|
||||
[HumanMessage(content="hello"), AIMessage(content="hi")]
|
||||
)
|
||||
|
||||
self.assertEqual(history, "User: hello\nAI: hi")
|
||||
|
||||
async def test_generate_response_streams_chunks(self):
|
||||
self.service.conversation_chain = FakeChain(chunks=["Hel", "lo!"])
|
||||
|
||||
chunks = [
|
||||
chunk
|
||||
async for chunk in self.service.generate_response(
|
||||
conversation(1), "hello", conversation_id=1
|
||||
)
|
||||
]
|
||||
|
||||
self.assertEqual("".join(chunks), "Hello!")
|
||||
|
||||
async def test_generate_response_sends_full_and_recent_history(self):
|
||||
self.service.conversation_chain = FakeChain(chunks=["ok"])
|
||||
messages = conversation(4) # 8 messages
|
||||
|
||||
async for _ in self.service.generate_response(messages, "latest", 1):
|
||||
pass
|
||||
|
||||
payload = self.service.conversation_chain.calls[0]
|
||||
self.assertEqual(payload["query"], "latest")
|
||||
self.assertEqual(len(payload["conversation"].splitlines()), 8)
|
||||
self.assertEqual(len(payload["recent_conversation"].splitlines()), 6)
|
||||
self.assertTrue(payload["recent_conversation"].endswith("AI: answer 3"))
|
||||
|
||||
|
||||
class SyncLLMServiceTestCase(SimpleTestCase):
|
||||
def test_generate_response_streams_chunks(self):
|
||||
service = SyncLLMService()
|
||||
service.conversation_chain = FakeChain(chunks=["one ", "two"])
|
||||
|
||||
chunks = list(service.generate_response(conversation=None, query="hello"))
|
||||
|
||||
self.assertEqual("".join(chunks), "one two")
|
||||
self.assertEqual(
|
||||
service.conversation_chain.calls, [{"query": "hello", "conversation": None}]
|
||||
)
|
||||
@@ -0,0 +1,279 @@
|
||||
import os
|
||||
import shutil
|
||||
import tempfile
|
||||
from unittest import mock
|
||||
|
||||
from asgiref.sync import sync_to_async
|
||||
from django.core.files.uploadedfile import SimpleUploadedFile
|
||||
from django.test import TransactionTestCase
|
||||
from langchain_community.document_loaders import (
|
||||
Docx2txtLoader,
|
||||
PyPDFLoader,
|
||||
TextLoader,
|
||||
UnstructuredFileLoader,
|
||||
)
|
||||
from langchain_core.documents import Document as LangDocument
|
||||
from langchain_core.messages import AIMessage, HumanMessage
|
||||
from parameterized import parameterized
|
||||
|
||||
from chat_backend.models import Document
|
||||
from chat_backend.services.rag_services import (
|
||||
AsyncRAGService,
|
||||
RAGService,
|
||||
SyncRAGService,
|
||||
get_documents,
|
||||
)
|
||||
|
||||
from .factories import make_company, make_workspace
|
||||
from .fakes import FakeChain
|
||||
|
||||
|
||||
def reset_singletons():
|
||||
for service in (RAGService, SyncRAGService, AsyncRAGService):
|
||||
service._instance = None
|
||||
|
||||
|
||||
class RAGServiceTestCase(TransactionTestCase):
|
||||
"""Exercises the RAG plumbing with Chroma and Ollama replaced by mocks.
|
||||
|
||||
``TransactionTestCase`` is required: the async helpers close old connections,
|
||||
which a class-wide atomic block (plain ``TestCase``) cannot survive.
|
||||
"""
|
||||
|
||||
def setUp(self):
|
||||
self.persist_dir = tempfile.mkdtemp()
|
||||
self.addCleanup(shutil.rmtree, self.persist_dir, True)
|
||||
|
||||
self.chroma = self._patch("chat_backend.services.rag_services.Chroma")
|
||||
# A fresh mock per call so re-initialisation is observable.
|
||||
self.chroma.side_effect = lambda *args, **kwargs: mock.MagicMock(
|
||||
name="vector_store"
|
||||
)
|
||||
self.embeddings = self._patch(
|
||||
"chat_backend.services.rag_services.OllamaEmbeddings"
|
||||
)
|
||||
self._patch("chat_backend.services.base_service.OllamaLLM")
|
||||
|
||||
overrides = self.settings(CHROMA_PERSIST_DIRECTORY=self.persist_dir)
|
||||
overrides.enable()
|
||||
self.addCleanup(overrides.disable)
|
||||
|
||||
reset_singletons()
|
||||
self.addCleanup(reset_singletons)
|
||||
self.service = AsyncRAGService()
|
||||
|
||||
self.workspace = make_workspace(make_company())
|
||||
|
||||
def _patch(self, target):
|
||||
patcher = mock.patch(target)
|
||||
self.addCleanup(patcher.stop)
|
||||
return patcher.start()
|
||||
|
||||
def _text_document(self, name="notes.txt", body=b"hello world"):
|
||||
return Document.objects.create(
|
||||
workspace=self.workspace, file=SimpleUploadedFile(name, body)
|
||||
)
|
||||
|
||||
def test_vector_store_uses_configured_persist_directory(self):
|
||||
self.assertTrue(os.path.isdir(self.persist_dir))
|
||||
self.chroma.assert_called_with(
|
||||
embedding_function=self.embeddings.return_value,
|
||||
persist_directory=self.persist_dir,
|
||||
)
|
||||
|
||||
@parameterized.expand(
|
||||
[
|
||||
("pdf", "/tmp/a.pdf", PyPDFLoader),
|
||||
("docx", "/tmp/a.docx", Docx2txtLoader),
|
||||
("txt_uppercase", "/tmp/A.TXT", TextLoader),
|
||||
("unknown", "/tmp/a.pptx", UnstructuredFileLoader),
|
||||
]
|
||||
)
|
||||
def test_loader_selection(self, _name, path, expected_loader):
|
||||
self.assertIs(self.service._get_file_loader(path), expected_loader)
|
||||
|
||||
@parameterized.expand(
|
||||
[
|
||||
("spaces_kept", "quarterly report.pdf", "quarterly report.pdf"),
|
||||
("parens", "report (final).pdf", "report _final_.pdf"),
|
||||
("path_separators", "../../etc/passwd", ".._.._etc_passwd"),
|
||||
]
|
||||
)
|
||||
def test_sanitize_filename(self, _name, given, expected):
|
||||
self.assertEqual(self.service._sanitize_filename(given), expected)
|
||||
|
||||
def test_materialize_file_field_writes_a_temp_copy(self):
|
||||
document = self._text_document(body=b"stored in postgres")
|
||||
|
||||
path = self.service._materialize_file_field(document.file)
|
||||
try:
|
||||
self.assertTrue(path.endswith(".txt"))
|
||||
with open(path, "rb") as handle:
|
||||
self.assertEqual(handle.read(), b"stored in postgres")
|
||||
finally:
|
||||
os.unlink(path)
|
||||
|
||||
def test_load_and_split_documents_merges_metadata(self):
|
||||
path = self._temp_text_file(b"sentence one. sentence two.")
|
||||
|
||||
chunks = self.service._load_and_split_documents(path, {"workspace_id": 7})
|
||||
|
||||
self.assertTrue(chunks)
|
||||
self.assertEqual(chunks[0].metadata["workspace_id"], 7)
|
||||
self.assertIn("sentence one", chunks[0].page_content)
|
||||
|
||||
def _temp_text_file(self, body: bytes) -> str:
|
||||
with tempfile.NamedTemporaryFile(suffix=".txt", delete=False) as handle:
|
||||
handle.write(body)
|
||||
self.addCleanup(lambda: os.path.exists(handle.name) and os.unlink(handle.name))
|
||||
return handle.name
|
||||
|
||||
def test_ingest_documents_adds_chunks_with_metadata(self):
|
||||
document = self._text_document(body=b"ingest me")
|
||||
|
||||
self.service.ingest_documents()
|
||||
|
||||
self.service.vector_store.add_documents.assert_called_once()
|
||||
added = self.service.vector_store.add_documents.call_args[0][0]
|
||||
self.assertIn("ingest me", added[0].page_content)
|
||||
self.assertEqual(added[0].metadata["workspace_id"], self.workspace.id)
|
||||
self.assertEqual(added[0].metadata["document_id"], document.id)
|
||||
|
||||
def test_ingest_documents_deletes_the_materialized_temp_file(self):
|
||||
self._text_document(body=b"ingest me")
|
||||
materialized = []
|
||||
original = self.service._materialize_file_field
|
||||
|
||||
def spy(file_field):
|
||||
path = original(file_field)
|
||||
materialized.append(path)
|
||||
return path
|
||||
|
||||
with mock.patch.object(
|
||||
self.service, "_materialize_file_field", side_effect=spy
|
||||
):
|
||||
self.service.ingest_documents()
|
||||
|
||||
self.assertTrue(materialized)
|
||||
self.assertFalse([path for path in materialized if os.path.exists(path)])
|
||||
|
||||
def test_ingest_documents_can_be_scoped_to_a_workspace(self):
|
||||
other_workspace = make_workspace(make_company("Other"), name="Other")
|
||||
Document.objects.create(
|
||||
workspace=other_workspace, file=SimpleUploadedFile("other.txt", b"other")
|
||||
)
|
||||
self._text_document(body=b"mine")
|
||||
|
||||
self.service.ingest_documents(workspace=self.workspace)
|
||||
|
||||
self.assertEqual(self.service.vector_store.add_documents.call_count, 1)
|
||||
added = self.service.vector_store.add_documents.call_args[0][0]
|
||||
self.assertIn("mine", added[0].page_content)
|
||||
|
||||
def test_add_files_to_store_reports_processed_files(self):
|
||||
path = self._temp_text_file(b"upload body")
|
||||
|
||||
results = self.service.add_files_to_store(
|
||||
[(path, "upload.txt", self.workspace.id)], workspace_id=self.workspace.id
|
||||
)
|
||||
|
||||
self.assertEqual(results["failed_files"], [])
|
||||
self.assertEqual(results["total_added"], 1)
|
||||
self.assertEqual(
|
||||
results["processed_files"],
|
||||
[{"filename": "upload.txt", "document_count": 1}],
|
||||
)
|
||||
added = self.service.vector_store.add_documents.call_args[0][0]
|
||||
self.assertEqual(added[0].metadata["source"], "upload.txt")
|
||||
self.assertEqual(added[0].metadata["workspace_id"], self.workspace.id)
|
||||
|
||||
def test_add_files_to_store_accepts_database_backed_file_fields(self):
|
||||
document = self._text_document(body=b"from the database")
|
||||
|
||||
results = self.service.add_files_to_store(
|
||||
[(document.file, document.file.name, self.workspace.id)],
|
||||
workspace_id=self.workspace.id,
|
||||
)
|
||||
|
||||
self.assertEqual(results["total_added"], 1)
|
||||
self.assertEqual(results["failed_files"], [])
|
||||
|
||||
def test_add_files_to_store_records_failures(self):
|
||||
results = self.service.add_files_to_store(
|
||||
[("/tmp/does-not-exist.txt", "missing.txt", self.workspace.id)],
|
||||
workspace_id=self.workspace.id,
|
||||
)
|
||||
|
||||
self.assertEqual(results["total_added"], 0)
|
||||
self.assertEqual(results["failed_files"][0]["filename"], "missing.txt")
|
||||
|
||||
def test_clear_vector_store_recreates_the_collection(self):
|
||||
first_store = self.service.vector_store
|
||||
|
||||
self.service.clear_vector_store()
|
||||
|
||||
first_store.delete_collection.assert_called_once()
|
||||
self.assertIsNot(self.service.vector_store, first_store)
|
||||
|
||||
async def test_search_documents_filters_by_workspace(self):
|
||||
retriever = self.service.vector_store.as_retriever.return_value
|
||||
retriever.aget_relevant_documents = mock.AsyncMock(
|
||||
return_value=[LangDocument(page_content="chunk")]
|
||||
)
|
||||
|
||||
docs = await self.service.search_documents("revenue", self.workspace, k=2)
|
||||
|
||||
self.assertEqual([doc.page_content for doc in docs], ["chunk"])
|
||||
self.service.vector_store.as_retriever.assert_called_with(
|
||||
search_type="mmr",
|
||||
search_kwargs={"k": 2, "filter": {"workspace_id": self.workspace.id}},
|
||||
)
|
||||
|
||||
async def test_search_documents_without_workspace_has_no_filter(self):
|
||||
retriever = self.service.vector_store.as_retriever.return_value
|
||||
retriever.aget_relevant_documents = mock.AsyncMock(return_value=[])
|
||||
|
||||
await self.service.search_documents("revenue")
|
||||
|
||||
self.service.vector_store.as_retriever.assert_called_with(
|
||||
search_type="mmr", search_kwargs={"k": 4, "filter": None}
|
||||
)
|
||||
|
||||
async def test_format_history_labels_speakers(self):
|
||||
history = await self.service._format_history(
|
||||
[
|
||||
HumanMessage(content="what is our policy?"),
|
||||
AIMessage(content="here it is"),
|
||||
]
|
||||
)
|
||||
|
||||
self.assertEqual(history, "User: what is our policy?\nAI: here it is")
|
||||
|
||||
async def test_generate_response_streams_chunks_with_history(self):
|
||||
self.service.rag_chain = FakeChain(chunks=["Based ", "on the docs."])
|
||||
conversation = [HumanMessage(content="what is our policy?")]
|
||||
|
||||
chunks = [
|
||||
chunk
|
||||
async for chunk in self.service.generate_response(
|
||||
conversation, "what is our policy?", self.workspace
|
||||
)
|
||||
]
|
||||
|
||||
self.assertEqual("".join(chunks), "Based on the docs.")
|
||||
payload = self.service.rag_chain.calls[0]
|
||||
self.assertEqual(payload["query"], "what is our policy?")
|
||||
self.assertEqual(payload["workspace"], self.workspace)
|
||||
self.assertEqual(payload["recent_conversation"], "User: what is our policy?")
|
||||
|
||||
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)(
|
||||
await sync_to_async(make_company)("Other"), name="Other"
|
||||
)
|
||||
|
||||
self.assertEqual([doc.id for doc in await get_documents()], [document.id])
|
||||
self.assertEqual(
|
||||
[doc.id for doc in await get_documents(self.workspace)], [document.id]
|
||||
)
|
||||
self.assertEqual(await get_documents(other_workspace), [])
|
||||
@@ -0,0 +1,54 @@
|
||||
import os
|
||||
from unittest import mock
|
||||
|
||||
from django.test import TestCase
|
||||
|
||||
from .factories import make_company, make_document, make_workspace
|
||||
|
||||
RAG_SERVICE = "chat_backend.services.rag_services.AsyncRAGService"
|
||||
|
||||
|
||||
class DocumentSignalTestCase(TestCase):
|
||||
def setUp(self):
|
||||
self.workspace = make_workspace(make_company())
|
||||
|
||||
def test_creating_a_document_reindexes_the_vector_store(self):
|
||||
with mock.patch.dict(os.environ, {"SKIP_RAG_INIT": ""}):
|
||||
with mock.patch(RAG_SERVICE) as service:
|
||||
make_document(self.workspace)
|
||||
|
||||
service.return_value.ingest_documents.assert_called_once_with()
|
||||
|
||||
def test_updating_a_document_does_not_reindex(self):
|
||||
document = make_document(self.workspace)
|
||||
|
||||
with mock.patch.dict(os.environ, {"SKIP_RAG_INIT": ""}):
|
||||
with mock.patch(RAG_SERVICE) as service:
|
||||
document.active = True
|
||||
document.save()
|
||||
|
||||
service.assert_not_called()
|
||||
|
||||
def test_deleting_a_document_reindexes_the_vector_store(self):
|
||||
document = make_document(self.workspace)
|
||||
|
||||
with mock.patch.dict(os.environ, {"SKIP_RAG_INIT": ""}):
|
||||
with mock.patch(RAG_SERVICE) as service:
|
||||
document.delete()
|
||||
|
||||
service.return_value.ingest_documents.assert_called_once_with()
|
||||
|
||||
def test_skip_rag_init_keeps_signals_inert(self):
|
||||
with mock.patch.dict(os.environ, {"SKIP_RAG_INIT": "1"}):
|
||||
with mock.patch(RAG_SERVICE) as service:
|
||||
document = make_document(self.workspace)
|
||||
document.delete()
|
||||
|
||||
service.assert_not_called()
|
||||
|
||||
def test_vector_store_failures_do_not_break_uploads(self):
|
||||
with mock.patch.dict(os.environ, {"SKIP_RAG_INIT": ""}):
|
||||
with mock.patch(RAG_SERVICE, side_effect=RuntimeError("chroma down")):
|
||||
document = make_document(self.workspace)
|
||||
|
||||
self.assertIsNotNone(document.pk)
|
||||
@@ -0,0 +1,94 @@
|
||||
from django.core.files.base import ContentFile
|
||||
from django.core.files.uploadedfile import SimpleUploadedFile
|
||||
from django.test import TestCase
|
||||
|
||||
from chat_backend.models import StoredFile
|
||||
from chat_backend.storage import DatabaseStorage
|
||||
|
||||
|
||||
class DatabaseStorageTestCase(TestCase):
|
||||
def setUp(self):
|
||||
self.storage = DatabaseStorage()
|
||||
|
||||
def test_save_then_open_roundtrip(self):
|
||||
name = self.storage.save("documents/hello.txt", ContentFile(b"hello world"))
|
||||
|
||||
self.assertEqual(name, "documents/hello.txt")
|
||||
with self.storage.open(name) as handle:
|
||||
self.assertEqual(handle.read(), b"hello world")
|
||||
|
||||
def test_save_records_size_and_content_type(self):
|
||||
name = self.storage.save(
|
||||
"documents/report.pdf",
|
||||
SimpleUploadedFile(
|
||||
"report.pdf", b"%PDF-1.3", content_type="application/pdf"
|
||||
),
|
||||
)
|
||||
|
||||
stored = StoredFile.objects.get(name=name)
|
||||
self.assertEqual(stored.size, len(b"%PDF-1.3"))
|
||||
self.assertEqual(stored.content_type, "application/pdf")
|
||||
self.assertEqual(self.storage.size(name), len(b"%PDF-1.3"))
|
||||
|
||||
def test_content_type_guessed_from_extension(self):
|
||||
name = self.storage.save("documents/notes.txt", ContentFile(b"abc"))
|
||||
self.assertEqual(StoredFile.objects.get(name=name).content_type, "text/plain")
|
||||
|
||||
def test_string_content_is_encoded(self):
|
||||
name = self.storage.save("documents/str.txt", ContentFile("héllo"))
|
||||
with self.storage.open(name) as handle:
|
||||
self.assertEqual(handle.read().decode("utf-8"), "héllo")
|
||||
|
||||
def test_second_save_of_same_name_gets_unique_name(self):
|
||||
first = self.storage.save("documents/dup.txt", ContentFile(b"one"))
|
||||
second = self.storage.save("documents/dup.txt", ContentFile(b"two"))
|
||||
|
||||
self.assertNotEqual(first, second)
|
||||
self.assertEqual(StoredFile.objects.count(), 2)
|
||||
with self.storage.open(first) as handle:
|
||||
self.assertEqual(handle.read(), b"one")
|
||||
|
||||
def test_exists_and_delete(self):
|
||||
name = self.storage.save("documents/gone.txt", ContentFile(b"bye"))
|
||||
|
||||
self.assertTrue(self.storage.exists(name))
|
||||
self.storage.delete(name)
|
||||
self.assertFalse(self.storage.exists(name))
|
||||
self.assertFalse(StoredFile.objects.filter(name=name).exists())
|
||||
|
||||
def test_delete_is_noop_for_missing_name(self):
|
||||
self.storage.delete("documents/never-written.txt")
|
||||
|
||||
def test_listdir_splits_dirs_and_files(self):
|
||||
self.storage.save("documents/top.txt", ContentFile(b"a"))
|
||||
self.storage.save("documents/nested/inner.txt", ContentFile(b"b"))
|
||||
|
||||
dirs, files = self.storage.listdir("documents")
|
||||
|
||||
self.assertEqual(files, ["top.txt"])
|
||||
self.assertEqual(dirs, ["nested"])
|
||||
|
||||
def test_url_points_at_api_route(self):
|
||||
self.assertEqual(
|
||||
self.storage.url("documents/a.txt"), "/api/stored-files/documents/a.txt"
|
||||
)
|
||||
|
||||
def test_created_and_modified_times(self):
|
||||
name = self.storage.save("documents/times.txt", ContentFile(b"a"))
|
||||
stored = StoredFile.objects.get(name=name)
|
||||
|
||||
self.assertEqual(self.storage.get_created_time(name), stored.created)
|
||||
self.assertEqual(self.storage.get_modified_time(name), stored.last_modified)
|
||||
|
||||
def test_path_and_accessed_time_are_unsupported(self):
|
||||
with self.assertRaises(NotImplementedError):
|
||||
self.storage.path("documents/a.txt")
|
||||
with self.assertRaises(NotImplementedError):
|
||||
self.storage.get_accessed_time("documents/a.txt")
|
||||
|
||||
def test_storage_is_deconstructible_for_migrations(self):
|
||||
path, args, kwargs = DatabaseStorage().deconstruct()
|
||||
|
||||
self.assertEqual(path, "chat_backend.storage.DatabaseStorage")
|
||||
self.assertEqual(list(args), [])
|
||||
self.assertEqual(dict(kwargs), {})
|
||||
@@ -0,0 +1,90 @@
|
||||
import datetime
|
||||
|
||||
from django.conf import settings
|
||||
from django.test import SimpleTestCase, override_settings
|
||||
from parameterized import parameterized
|
||||
|
||||
from chat_backend.ollama_config import (
|
||||
ollama_base_url,
|
||||
ollama_embed_model,
|
||||
ollama_embeddings_kwargs,
|
||||
ollama_llm_kwargs,
|
||||
ollama_model,
|
||||
)
|
||||
from chat_backend.utils import last_day_of_month
|
||||
|
||||
|
||||
class LastDayOfMonthTestCase(SimpleTestCase):
|
||||
@parameterized.expand(
|
||||
[
|
||||
("january", datetime.date(2025, 1, 5), datetime.date(2025, 1, 31)),
|
||||
("leap_february", datetime.date(2024, 2, 10), datetime.date(2024, 2, 29)),
|
||||
("common_february", datetime.date(2025, 2, 10), datetime.date(2025, 2, 28)),
|
||||
("april", datetime.date(2025, 4, 1), datetime.date(2025, 4, 30)),
|
||||
("december", datetime.date(2025, 12, 31), datetime.date(2025, 12, 31)),
|
||||
]
|
||||
)
|
||||
def test_returns_last_day(self, _name, given, expected):
|
||||
self.assertEqual(last_day_of_month(given), expected)
|
||||
|
||||
def test_keeps_datetime_time_component(self):
|
||||
result = last_day_of_month(datetime.datetime(2025, 6, 4, 13, 45))
|
||||
|
||||
self.assertEqual(result.date(), datetime.date(2025, 6, 30))
|
||||
self.assertEqual((result.hour, result.minute), (13, 45))
|
||||
|
||||
|
||||
@override_settings(
|
||||
OLLAMA_BASE_URL="http://10.0.0.128:11434",
|
||||
OLLAMA_MODEL="llama3.2",
|
||||
OLLAMA_EMBED_MODEL="nomic-embed-text",
|
||||
)
|
||||
class OllamaConfigTestCase(SimpleTestCase):
|
||||
def test_reads_settings(self):
|
||||
self.assertEqual(ollama_base_url(), "http://10.0.0.128:11434")
|
||||
self.assertEqual(ollama_model(), "llama3.2")
|
||||
self.assertEqual(ollama_embed_model(), "nomic-embed-text")
|
||||
|
||||
def test_model_argument_wins_over_settings(self):
|
||||
self.assertEqual(ollama_model("gpt-oss:20b"), "gpt-oss:20b")
|
||||
|
||||
def test_llm_kwargs_carry_base_url_and_model(self):
|
||||
kwargs = ollama_llm_kwargs(temperature=0.1, num_ctx=2048)
|
||||
|
||||
self.assertEqual(
|
||||
kwargs,
|
||||
{
|
||||
"base_url": "http://10.0.0.128:11434",
|
||||
"model": "llama3.2",
|
||||
"temperature": 0.1,
|
||||
"num_ctx": 2048,
|
||||
},
|
||||
)
|
||||
|
||||
def test_llm_kwargs_extra_overrides_defaults(self):
|
||||
self.assertEqual(ollama_llm_kwargs(model="other")["model"], "other")
|
||||
|
||||
def test_embeddings_kwargs_use_embed_model(self):
|
||||
self.assertEqual(
|
||||
ollama_embeddings_kwargs(),
|
||||
{"base_url": "http://10.0.0.128:11434", "model": "nomic-embed-text"},
|
||||
)
|
||||
|
||||
|
||||
class OllamaConfigFallbackTestCase(SimpleTestCase):
|
||||
"""Deployed code must never fall back to a hardcoded remote host."""
|
||||
|
||||
def test_defaults_when_settings_are_absent(self):
|
||||
with override_settings():
|
||||
del settings.OLLAMA_BASE_URL
|
||||
del settings.OLLAMA_MODEL
|
||||
del settings.OLLAMA_EMBED_MODEL
|
||||
|
||||
self.assertEqual(ollama_base_url(), "http://127.0.0.1:11434")
|
||||
self.assertEqual(ollama_model(), "llama3.2")
|
||||
self.assertEqual(ollama_embed_model(), "llama3.2")
|
||||
|
||||
def test_embed_model_falls_back_to_chat_model(self):
|
||||
with override_settings(OLLAMA_MODEL="llama3.2"):
|
||||
del settings.OLLAMA_EMBED_MODEL
|
||||
self.assertEqual(ollama_embed_model(), "llama3.2")
|
||||
@@ -0,0 +1,129 @@
|
||||
import datetime
|
||||
|
||||
from django.urls import reverse
|
||||
from django.utils import timezone
|
||||
from rest_framework import status
|
||||
from rest_framework.test import APITestCase
|
||||
|
||||
from chat_backend.models import PromptMetric
|
||||
|
||||
from .factories import make_company, make_conversation, make_prompt, make_user
|
||||
|
||||
|
||||
def mid_month():
|
||||
"""The 2nd of the current month — inside the window the views query, whatever today is."""
|
||||
return timezone.now().replace(day=2, hour=12, minute=0, second=0, microsecond=0)
|
||||
|
||||
|
||||
class AnalyticsTestData(APITestCase):
|
||||
"""Two users in one company, one user in another, all active this month."""
|
||||
|
||||
def setUp(self):
|
||||
self.company = make_company()
|
||||
self.user = make_user(email="me@example.com", company=self.company)
|
||||
self.teammate = make_user(email="mate@example.com", company=self.company)
|
||||
self.outsider = make_user(
|
||||
email="out@example.com", company=make_company("Other")
|
||||
)
|
||||
self.client.force_authenticate(user=self.user)
|
||||
|
||||
self.my_conversation = self._conversation(self.user, prompts=2)
|
||||
self._conversation(self.teammate, prompts=1)
|
||||
self._conversation(self.outsider, prompts=3)
|
||||
|
||||
def _conversation(self, user, prompts: int):
|
||||
conversation = make_conversation(user=user, created=mid_month())
|
||||
for index in range(prompts):
|
||||
make_prompt(conversation, message=f"prompt {index}", created=mid_month())
|
||||
return conversation
|
||||
|
||||
|
||||
class UserPromptAnalyticsTestCase(AnalyticsTestData):
|
||||
def test_returns_three_months_oldest_first(self):
|
||||
response = self.client.get(reverse("analytics_user_prompts"))
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
self.assertEqual(len(response.data), 3)
|
||||
self.assertEqual(response.data[-1]["month"], timezone.now().strftime("%B"))
|
||||
|
||||
def test_current_month_counts(self):
|
||||
current = self.client.get(reverse("analytics_user_prompts")).data[-1]
|
||||
|
||||
self.assertEqual(current["you"], 2)
|
||||
self.assertEqual(current["others"], 3 / 2)
|
||||
self.assertEqual(current["all"], 6 / 3)
|
||||
|
||||
def test_earlier_months_are_empty(self):
|
||||
earlier = self.client.get(reverse("analytics_user_prompts")).data[0]
|
||||
|
||||
self.assertEqual(earlier["you"], 0)
|
||||
self.assertEqual(earlier["others"], 0)
|
||||
self.assertEqual(earlier["all"], 0)
|
||||
|
||||
|
||||
class UserConversationAnalyticsTestCase(AnalyticsTestData):
|
||||
def test_current_month_counts(self):
|
||||
response = self.client.get(reverse("analytics_user_conversations"))
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
current = response.data[-1]
|
||||
self.assertEqual(current["you"], 1)
|
||||
self.assertEqual(current["others"], 2 / 2)
|
||||
self.assertEqual(current["all"], 3 / 3)
|
||||
|
||||
|
||||
class CompanyUsageAnalyticsTestCase(AnalyticsTestData):
|
||||
def test_counts_active_and_idle_company_users(self):
|
||||
response = self.client.get(reverse("analytics_company_usage"))
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
current = response.data[-1]
|
||||
self.assertEqual(current["used"], 2)
|
||||
self.assertEqual(current["not_used"], 0)
|
||||
|
||||
def test_idle_users_are_reported(self):
|
||||
make_user(email="idle@example.com", company=self.company)
|
||||
|
||||
current = self.client.get(reverse("analytics_company_usage")).data[-1]
|
||||
|
||||
self.assertEqual(current["used"], 2)
|
||||
self.assertEqual(current["not_used"], 1)
|
||||
|
||||
|
||||
class AdminAnalyticsTestCase(APITestCase):
|
||||
def setUp(self):
|
||||
self.user = make_user(company=make_company())
|
||||
self.client.force_authenticate(user=self.user)
|
||||
|
||||
def _metric(self, seconds: int):
|
||||
start = mid_month()
|
||||
return PromptMetric.objects.create(
|
||||
prompt_id=1,
|
||||
conversation_id=1,
|
||||
model_name="llama3.2",
|
||||
start_time=start,
|
||||
end_time=start + datetime.timedelta(seconds=seconds),
|
||||
prompt_length=10,
|
||||
has_file=False,
|
||||
created=start,
|
||||
)
|
||||
|
||||
def test_zeros_when_no_metrics(self):
|
||||
response = self.client.get(reverse("analytics_admin"))
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
self.assertEqual(len(response.data), 3)
|
||||
for month in response.data:
|
||||
self.assertEqual(month["range"], [0, 0])
|
||||
self.assertEqual(month["avg"], 0)
|
||||
self.assertEqual(month["median"], 0)
|
||||
|
||||
def test_duration_statistics_for_current_month(self):
|
||||
self._metric(10)
|
||||
self._metric(20)
|
||||
|
||||
current = self.client.get(reverse("analytics_admin")).data[-1]
|
||||
|
||||
self.assertEqual(current["range"], [10, 20])
|
||||
self.assertEqual(current["avg"], 15)
|
||||
self.assertEqual(current["median"], 20)
|
||||
@@ -0,0 +1,141 @@
|
||||
from django.urls import reverse
|
||||
from rest_framework import status
|
||||
from rest_framework.test import APITestCase
|
||||
|
||||
from chat_backend.models import Conversation, Prompt
|
||||
|
||||
from .factories import make_company, make_conversation, make_prompt, make_user
|
||||
|
||||
|
||||
class ConversationsViewTestCase(APITestCase):
|
||||
def setUp(self):
|
||||
self.user = make_user(company=make_company())
|
||||
self.client.force_authenticate(user=self.user)
|
||||
self.url = reverse("conversations")
|
||||
|
||||
def test_list_hides_deleted_and_other_users_conversations(self):
|
||||
make_conversation(user=self.user, title="Mine")
|
||||
make_conversation(user=self.user, title="Removed", deleted=True)
|
||||
make_conversation(user=make_user(email="other@example.com"), title="Theirs")
|
||||
|
||||
response = self.client.get(self.url)
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
self.assertEqual([item["title"] for item in response.data], ["Mine"])
|
||||
|
||||
def test_list_order_follows_user_preference(self):
|
||||
make_conversation(user=self.user, title="First")
|
||||
make_conversation(user=self.user, title="Second")
|
||||
|
||||
oldest_first = self.client.get(self.url).data
|
||||
self.assertEqual([item["title"] for item in oldest_first], ["First", "Second"])
|
||||
|
||||
self.user.conversation_order = False
|
||||
self.user.save()
|
||||
|
||||
newest_first = self.client.get(self.url).data
|
||||
self.assertEqual([item["title"] for item in newest_first], ["Second", "First"])
|
||||
|
||||
def test_post_creates_conversation_for_current_user(self):
|
||||
response = self.client.post(self.url, {"name": "New Chat"}, format="json")
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_201_CREATED)
|
||||
conversation = Conversation.objects.get(id=response.data["id"])
|
||||
self.assertEqual(conversation.title, "New Chat")
|
||||
self.assertEqual(conversation.user, self.user)
|
||||
self.assertEqual(response.data["title"], "New Chat")
|
||||
|
||||
|
||||
class ConversationPreferencesTestCase(APITestCase):
|
||||
def setUp(self):
|
||||
self.user = make_user(company=make_company())
|
||||
self.client.force_authenticate(user=self.user)
|
||||
self.url = reverse("conversation_preferences")
|
||||
|
||||
def test_get_returns_current_order(self):
|
||||
response = self.client.get(self.url)
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
self.assertTrue(response.data["order"])
|
||||
|
||||
def test_post_toggles_and_persists_order(self):
|
||||
response = self.client.post(self.url)
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
self.assertFalse(response.data["order"])
|
||||
self.user.refresh_from_db()
|
||||
self.assertFalse(self.user.conversation_order)
|
||||
|
||||
|
||||
class ConversationDetailViewTestCase(APITestCase):
|
||||
def setUp(self):
|
||||
self.user = make_user(company=make_company())
|
||||
self.client.force_authenticate(user=self.user)
|
||||
self.conversation = make_conversation(user=self.user)
|
||||
self.url = reverse("conversation_details")
|
||||
|
||||
def test_get_returns_prompts_of_conversation(self):
|
||||
make_prompt(self.conversation, message="hello")
|
||||
make_prompt(self.conversation, message="hi there", user_created=False)
|
||||
make_prompt(make_conversation(user=self.user), message="other conversation")
|
||||
|
||||
response = self.client.get(self.url, {"conversation_id": self.conversation.id})
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
self.assertEqual(
|
||||
[(item["message"], item["user_created"]) for item in response.data],
|
||||
[("hello", True), ("hi there", False)],
|
||||
)
|
||||
|
||||
def test_post_stores_assistant_prompt(self):
|
||||
response = self.client.post(
|
||||
self.url,
|
||||
{
|
||||
"prompt": "generated answer",
|
||||
"conversation_id": self.conversation.id,
|
||||
"is_user": False,
|
||||
},
|
||||
format="json",
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
prompt = Prompt.objects.get()
|
||||
self.assertEqual(prompt.message, "generated answer")
|
||||
self.assertFalse(prompt.user_created)
|
||||
self.assertEqual(prompt.conversation_id, self.conversation.id)
|
||||
|
||||
def test_post_stores_user_prompt_and_broadcasts(self):
|
||||
response = self.client.post(
|
||||
self.url,
|
||||
{
|
||||
"prompt": "what is the weather",
|
||||
"conversation_id": self.conversation.id,
|
||||
"is_user": True,
|
||||
},
|
||||
format="json",
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
prompt = Prompt.objects.get()
|
||||
self.assertTrue(prompt.user_created)
|
||||
self.assertEqual(prompt.conversation_id, self.conversation.id)
|
||||
|
||||
def test_post_for_unknown_conversation_stores_nothing(self):
|
||||
response = self.client.post(
|
||||
self.url,
|
||||
{"prompt": "orphan", "conversation_id": 4242, "is_user": True},
|
||||
format="json",
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
self.assertEqual(Prompt.objects.count(), 0)
|
||||
|
||||
def test_delete_soft_deletes_conversation(self):
|
||||
response = self.client.delete(
|
||||
self.url, {"conversation_id": self.conversation.id}, format="json"
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_202_ACCEPTED)
|
||||
self.conversation.refresh_from_db()
|
||||
self.assertTrue(self.conversation.deleted)
|
||||
self.assertEqual(Conversation.objects.count(), 1)
|
||||
@@ -0,0 +1,142 @@
|
||||
from unittest import mock
|
||||
|
||||
from django.urls import reverse
|
||||
from rest_framework import status
|
||||
from rest_framework.test import APITestCase
|
||||
|
||||
from chat_backend.models import Document, DocumentWorkspace, StoredFile
|
||||
|
||||
from .factories import (
|
||||
make_company,
|
||||
make_document,
|
||||
make_user,
|
||||
make_workspace,
|
||||
pdf_upload,
|
||||
)
|
||||
|
||||
|
||||
class DocumentWorkspaceViewTestCase(APITestCase):
|
||||
def setUp(self):
|
||||
self.company = make_company()
|
||||
self.user = make_user(company=self.company)
|
||||
self.client.force_authenticate(user=self.user)
|
||||
self.workspace = make_workspace(self.company)
|
||||
self.url = reverse("document_workspaces")
|
||||
|
||||
def test_list_only_returns_own_company_workspaces(self):
|
||||
make_workspace(make_company("Other"), name="Other Workspace")
|
||||
|
||||
response = self.client.get(self.url)
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
self.assertEqual([item["name"] for item in response.data], ["Test Workspace"])
|
||||
|
||||
def test_create_workspace_attaches_company(self):
|
||||
response = self.client.post(self.url, {"name": "New Workspace"}, format="json")
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_201_CREATED)
|
||||
self.assertEqual(DocumentWorkspace.objects.count(), 2)
|
||||
created = DocumentWorkspace.objects.get(name="New Workspace")
|
||||
self.assertEqual(created.company, self.company)
|
||||
|
||||
def test_create_workspace_requires_name(self):
|
||||
response = self.client.post(self.url, {}, format="json")
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST)
|
||||
self.assertIn("name", response.data)
|
||||
|
||||
|
||||
class DocumentUploadViewTestCase(APITestCase):
|
||||
def setUp(self):
|
||||
self.company = make_company()
|
||||
self.user = make_user(company=self.company)
|
||||
self.client.force_authenticate(user=self.user)
|
||||
self.workspace = make_workspace(self.company)
|
||||
self.url = reverse("documents")
|
||||
rag_patcher = mock.patch("chat_backend.views.AsyncRAGService")
|
||||
self.rag_service = rag_patcher.start()
|
||||
self.addCleanup(rag_patcher.stop)
|
||||
|
||||
def test_upload_stores_bytes_in_database_and_marks_processed(self):
|
||||
response = self.client.post(
|
||||
self.url, {"file": pdf_upload()}, format="multipart"
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_201_CREATED)
|
||||
document = Document.objects.get()
|
||||
self.assertEqual(document.workspace, self.workspace)
|
||||
self.assertTrue(document.processed)
|
||||
self.assertTrue(document.active)
|
||||
self.assertTrue(StoredFile.objects.filter(name=document.file.name).exists())
|
||||
|
||||
def test_upload_hands_file_to_the_rag_service(self):
|
||||
self.client.post(self.url, {"file": pdf_upload()}, format="multipart")
|
||||
|
||||
self.rag_service.return_value.add_files_to_store.assert_called_once()
|
||||
_args, kwargs = self.rag_service.return_value.add_files_to_store.call_args
|
||||
self.assertEqual(kwargs["workspace_id"], self.workspace.id)
|
||||
|
||||
def test_upload_without_file_is_rejected(self):
|
||||
response = self.client.post(self.url, {}, format="multipart")
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST)
|
||||
self.assertEqual(Document.objects.count(), 0)
|
||||
|
||||
def test_upload_of_non_file_payload_is_rejected(self):
|
||||
response = self.client.post(
|
||||
self.url, {"file": "not a file"}, format="multipart"
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST)
|
||||
|
||||
def test_upload_without_workspace_returns_404(self):
|
||||
self.workspace.delete()
|
||||
|
||||
response = self.client.post(
|
||||
self.url, {"file": pdf_upload()}, format="multipart"
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_404_NOT_FOUND)
|
||||
|
||||
def test_list_returns_documents_of_own_workspace(self):
|
||||
make_document(self.workspace)
|
||||
|
||||
response = self.client.get(self.url)
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
self.assertEqual(len(response.data), 1)
|
||||
self.assertIn("test", response.data[0]["file"])
|
||||
self.assertIn("pdf", response.data[0]["file"])
|
||||
|
||||
def test_list_without_workspace_returns_404(self):
|
||||
self.workspace.delete()
|
||||
|
||||
response = self.client.get(self.url)
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_404_NOT_FOUND)
|
||||
self.assertEqual(response.data["error"], "Workspace not found")
|
||||
|
||||
|
||||
class DocumentDetailViewTestCase(APITestCase):
|
||||
def setUp(self):
|
||||
self.company = make_company()
|
||||
self.user = make_user(company=self.company)
|
||||
self.client.force_authenticate(user=self.user)
|
||||
self.workspace = make_workspace(self.company)
|
||||
|
||||
def test_other_companies_documents_are_not_reachable(self):
|
||||
other_workspace = make_workspace(make_company("Other"), name="Other Workspace")
|
||||
other_document = make_document(other_workspace)
|
||||
|
||||
url = reverse("documents_details", kwargs={"document_id": other_document.id})
|
||||
response = self.client.get(url)
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_404_NOT_FOUND)
|
||||
|
||||
def test_unknown_document_returns_404(self):
|
||||
url = reverse("documents_details", kwargs={"document_id": 4242})
|
||||
|
||||
response = self.client.get(url)
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_404_NOT_FOUND)
|
||||
self.assertEqual(response.data["error"], "Document not found")
|
||||
@@ -0,0 +1,354 @@
|
||||
from django.core import mail
|
||||
from django.urls import reverse
|
||||
from rest_framework import status
|
||||
from rest_framework.test import APITestCase
|
||||
from rest_framework_simplejwt.tokens import RefreshToken
|
||||
|
||||
from chat_backend.models import Announcement, CustomUser, Feedback
|
||||
|
||||
from .factories import make_company, make_user
|
||||
|
||||
|
||||
class AuthenticationRequiredTestCase(APITestCase):
|
||||
def test_protected_endpoints_reject_anonymous_requests(self):
|
||||
for name in [
|
||||
"get_user",
|
||||
"conversations",
|
||||
"feedbacks",
|
||||
"company_users",
|
||||
"conversation_preferences",
|
||||
"analytics_user_prompts",
|
||||
"documents",
|
||||
]:
|
||||
with self.subTest(endpoint=name):
|
||||
response = self.client.get(reverse(name))
|
||||
self.assertEqual(response.status_code, status.HTTP_401_UNAUTHORIZED)
|
||||
|
||||
def test_announcements_are_public(self):
|
||||
Announcement.objects.create(message="scheduled maintenance")
|
||||
|
||||
response = self.client.get(reverse("get_announcments"))
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
self.assertEqual(response.data[0]["message"], "scheduled maintenance")
|
||||
|
||||
|
||||
class TokenTestCase(APITestCase):
|
||||
def setUp(self):
|
||||
self.user = make_user(email="person@example.com", password="testpass123")
|
||||
|
||||
def test_obtain_token_pair(self):
|
||||
response = self.client.post(
|
||||
reverse("token_create"),
|
||||
{"username": self.user.username, "password": "testpass123"},
|
||||
format="json",
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
self.assertIn("access", response.data)
|
||||
self.assertIn("refresh", response.data)
|
||||
|
||||
def test_obtain_token_rejects_bad_password(self):
|
||||
response = self.client.post(
|
||||
reverse("token_create"),
|
||||
{"username": self.user.username, "password": "wrong"},
|
||||
format="json",
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_401_UNAUTHORIZED)
|
||||
|
||||
def test_access_token_authenticates_requests(self):
|
||||
token = self.client.post(
|
||||
reverse("token_create"),
|
||||
{"username": self.user.username, "password": "testpass123"},
|
||||
format="json",
|
||||
).data["access"]
|
||||
|
||||
self.client.credentials(HTTP_AUTHORIZATION=f"JWT {token}")
|
||||
response = self.client.get(reverse("get_user"))
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
self.assertEqual(response.data["email"], self.user.email)
|
||||
|
||||
def test_blacklist_refresh_token(self):
|
||||
refresh = RefreshToken.for_user(self.user)
|
||||
|
||||
response = self.client.post(
|
||||
reverse("blacklist"), {"refresh_token": str(refresh)}, format="json"
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_205_RESET_CONTENT)
|
||||
|
||||
def test_blacklist_rejects_garbage_token(self):
|
||||
response = self.client.post(
|
||||
reverse("blacklist"), {"refresh_token": "not-a-token"}, format="json"
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST)
|
||||
|
||||
def test_is_authenticated_endpoint(self):
|
||||
url = reverse("is_authenticated")
|
||||
|
||||
self.assertEqual(self.client.get(url).status_code, status.HTTP_401_UNAUTHORIZED)
|
||||
|
||||
self.client.force_login(self.user)
|
||||
self.assertEqual(self.client.get(url).status_code, status.HTTP_200_OK)
|
||||
|
||||
|
||||
class UserCreateTestCase(APITestCase):
|
||||
def test_invalid_payload_returns_errors(self):
|
||||
response = self.client.post(reverse("create_user"), {}, format="json")
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST)
|
||||
self.assertIn("email", response.data)
|
||||
self.assertEqual(CustomUser.objects.count(), 0)
|
||||
|
||||
|
||||
class SetPasswordTestCase(APITestCase):
|
||||
def setUp(self):
|
||||
self.user = make_user(email="person@example.com", password="testpass123")
|
||||
|
||||
def test_get_rejects_user_that_already_has_a_password(self):
|
||||
url = reverse("set_password", kwargs={"slug": self.user.slug})
|
||||
|
||||
self.assertEqual(self.client.get(url).status_code, status.HTTP_401_UNAUTHORIZED)
|
||||
|
||||
def test_get_allows_user_awaiting_a_password(self):
|
||||
self.user.set_unusable_password()
|
||||
self.user.save()
|
||||
|
||||
url = reverse("set_password", kwargs={"slug": self.user.slug})
|
||||
|
||||
self.assertEqual(self.client.get(url).status_code, status.HTTP_200_OK)
|
||||
|
||||
def test_post_sets_password(self):
|
||||
url = reverse("set_password", kwargs={"slug": self.user.slug})
|
||||
|
||||
response = self.client.post(url, {"password": "brandnewpass"}, format="json")
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
self.user.refresh_from_db()
|
||||
self.assertTrue(self.user.check_password("brandnewpass"))
|
||||
|
||||
|
||||
class AcknowledgeTermsOfServiceTestCase(APITestCase):
|
||||
def test_post_marks_tos_signed(self):
|
||||
user = make_user()
|
||||
self.client.force_authenticate(user=user)
|
||||
|
||||
response = self.client.post(reverse("acknowledge_tos"))
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
user.refresh_from_db()
|
||||
self.assertTrue(user.has_signed_tos)
|
||||
|
||||
|
||||
class UserInviteTestCase(APITestCase):
|
||||
def setUp(self):
|
||||
self.company = make_company()
|
||||
self.manager = make_user(
|
||||
email="manager@example.com", company=self.company, is_company_manager=True
|
||||
)
|
||||
self.url = reverse("invite_user")
|
||||
|
||||
def test_manager_invites_new_user_and_email_is_sent(self):
|
||||
self.client.force_authenticate(user=self.manager)
|
||||
|
||||
response = self.client.post(
|
||||
self.url, {"email": "newhire@example.com"}, format="json"
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_201_CREATED)
|
||||
invited = CustomUser.objects.get(email="newhire@example.com")
|
||||
self.assertEqual(invited.company_id, self.company.id)
|
||||
self.assertEqual(invited.username, "newhire@example.com")
|
||||
self.assertEqual(len(mail.outbox), 1)
|
||||
self.assertIn(invited.slug, mail.outbox[0].body)
|
||||
self.assertEqual(mail.outbox[0].to, ["newhire@example.com"])
|
||||
|
||||
def test_non_manager_cannot_invite(self):
|
||||
member = make_user(email="member@example.com", company=self.company)
|
||||
self.client.force_authenticate(user=member)
|
||||
|
||||
response = self.client.post(
|
||||
self.url, {"email": "newhire@example.com"}, format="json"
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST)
|
||||
self.assertFalse(
|
||||
CustomUser.objects.filter(email="newhire@example.com").exists()
|
||||
)
|
||||
|
||||
def test_malformed_email_is_rejected(self):
|
||||
self.client.force_authenticate(user=self.manager)
|
||||
|
||||
response = self.client.post(self.url, {"email": "nope"}, format="json")
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST)
|
||||
self.assertEqual(len(mail.outbox), 0)
|
||||
|
||||
def test_duplicate_email_is_rejected(self):
|
||||
self.client.force_authenticate(user=self.manager)
|
||||
|
||||
response = self.client.post(
|
||||
self.url, {"email": self.manager.email}, format="json"
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST)
|
||||
|
||||
def test_get_is_not_allowed(self):
|
||||
self.client.force_authenticate(user=self.manager)
|
||||
|
||||
response = self.client.get(self.url)
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_405_METHOD_NOT_ALLOWED)
|
||||
|
||||
|
||||
class FeedbackViewTestCase(APITestCase):
|
||||
def setUp(self):
|
||||
self.user = make_user(company=make_company())
|
||||
self.client.force_authenticate(user=self.user)
|
||||
self.url = reverse("feedbacks")
|
||||
|
||||
def test_post_creates_feedback_and_notifies(self):
|
||||
response = self.client.post(
|
||||
self.url,
|
||||
{"title": "Broken button", "text": "It does nothing"},
|
||||
format="json",
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_201_CREATED)
|
||||
feedback = Feedback.objects.get()
|
||||
self.assertEqual(feedback.user, self.user)
|
||||
self.assertEqual(len(mail.outbox), 1)
|
||||
self.assertIn("Broken button", mail.outbox[0].body)
|
||||
|
||||
def test_post_without_text_is_rejected(self):
|
||||
response = self.client.post(self.url, {"title": "no body"}, format="json")
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST)
|
||||
self.assertEqual(Feedback.objects.count(), 0)
|
||||
|
||||
def test_get_only_returns_own_feedback(self):
|
||||
Feedback.objects.create(title="mine", text="a", user=self.user)
|
||||
Feedback.objects.create(
|
||||
title="theirs", text="b", user=make_user(email="o@o.com")
|
||||
)
|
||||
|
||||
response = self.client.get(self.url)
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
self.assertEqual([item["title"] for item in response.data], ["mine"])
|
||||
|
||||
|
||||
class CompanyUsersViewTestCase(APITestCase):
|
||||
def setUp(self):
|
||||
self.company = make_company()
|
||||
self.manager = make_user(
|
||||
email="manager@example.com", company=self.company, is_company_manager=True
|
||||
)
|
||||
self.member = make_user(email="member@example.com", company=self.company)
|
||||
self.url = reverse("company_users")
|
||||
|
||||
def test_manager_lists_company_users(self):
|
||||
make_user(email="outsider@example.com", company=make_company("Other"))
|
||||
self.client.force_authenticate(user=self.manager)
|
||||
|
||||
response = self.client.get(self.url)
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
self.assertEqual(
|
||||
{item["email"] for item in response.data},
|
||||
{"manager@example.com", "member@example.com"},
|
||||
)
|
||||
|
||||
def test_member_cannot_list_company_users(self):
|
||||
self.client.force_authenticate(user=self.member)
|
||||
|
||||
response = self.client.get(self.url)
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_401_UNAUTHORIZED)
|
||||
|
||||
def test_manager_toggles_is_active(self):
|
||||
self.client.force_authenticate(user=self.manager)
|
||||
|
||||
response = self.client.post(
|
||||
self.url, {"email": self.member.email, "field": "is_active"}, format="json"
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
self.member.refresh_from_db()
|
||||
self.assertFalse(self.member.is_active)
|
||||
|
||||
def test_manager_toggles_company_manager(self):
|
||||
self.client.force_authenticate(user=self.manager)
|
||||
|
||||
response = self.client.post(
|
||||
self.url,
|
||||
{"email": self.member.email, "field": "company_manager"},
|
||||
format="json",
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
self.member.refresh_from_db()
|
||||
self.assertTrue(self.member.is_company_manager)
|
||||
|
||||
def test_manager_clears_password(self):
|
||||
self.client.force_authenticate(user=self.manager)
|
||||
|
||||
response = self.client.post(
|
||||
self.url,
|
||||
{"email": self.member.email, "field": "has_password"},
|
||||
format="json",
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
self.member.refresh_from_db()
|
||||
self.assertFalse(self.member.has_usable_password())
|
||||
|
||||
def test_manager_cannot_touch_other_company_users(self):
|
||||
outsider = make_user(
|
||||
email="outsider@example.com", company=make_company("Other")
|
||||
)
|
||||
self.client.force_authenticate(user=self.manager)
|
||||
|
||||
response = self.client.post(
|
||||
self.url, {"email": outsider.email, "field": "is_active"}, format="json"
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_401_UNAUTHORIZED)
|
||||
outsider.refresh_from_db()
|
||||
self.assertTrue(outsider.is_active)
|
||||
|
||||
def test_manager_deletes_company_user(self):
|
||||
self.client.force_authenticate(user=self.manager)
|
||||
|
||||
response = self.client.delete(
|
||||
self.url, {"email": self.member.email}, format="json"
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
self.assertFalse(CustomUser.objects.filter(email=self.member.email).exists())
|
||||
|
||||
def test_member_cannot_delete_users(self):
|
||||
self.client.force_authenticate(user=self.member)
|
||||
|
||||
response = self.client.delete(
|
||||
self.url, {"email": self.manager.email}, format="json"
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_401_UNAUTHORIZED)
|
||||
self.assertTrue(CustomUser.objects.filter(email=self.manager.email).exists())
|
||||
|
||||
|
||||
class CustomUserGetTestCase(APITestCase):
|
||||
def test_returns_authenticated_user(self):
|
||||
user = make_user(company=make_company("Globex"))
|
||||
self.client.force_authenticate(user=user)
|
||||
|
||||
response = self.client.get(reverse("get_user"))
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
self.assertEqual(response.data["email"], user.email)
|
||||
self.assertEqual(response.data["company"]["name"], "Globex")
|
||||
self.assertNotIn("password", response.data)
|
||||
@@ -30,7 +30,6 @@ from .models import (
|
||||
)
|
||||
from django.views.decorators.cache import never_cache
|
||||
from django.http import JsonResponse
|
||||
from datetime import datetime
|
||||
from .client import LlamaClient
|
||||
from asgiref.sync import sync_to_async, async_to_sync
|
||||
from channels.generic.websocket import AsyncWebsocketConsumer
|
||||
@@ -73,19 +72,19 @@ 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__)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
CHANNEL_NAME: str = "llm_messages"
|
||||
MODEL_NAME: str = ollama_model()
|
||||
|
||||
|
||||
# Create your views here.
|
||||
class CustomObtainTokenView(TokenObtainPairView):
|
||||
permission_classes = (permissions.AllowAny,)
|
||||
@@ -480,7 +479,7 @@ class ConversationDetailView(APIView):
|
||||
data={
|
||||
"message": prompt,
|
||||
"user_created": is_user,
|
||||
"created": datetime.now(),
|
||||
"created": timezone.now(),
|
||||
}
|
||||
)
|
||||
if serializer.is_valid():
|
||||
@@ -725,9 +724,6 @@ llm = OllamaLLM(**ollama_llm_kwargs(model=MODEL_NAME))
|
||||
# chain = prompt | llm.with_config({"run_name": "model"}) | output_parser.with_config({"run_name": "Assistant"})
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# Document Views
|
||||
class DocumentWorkspaceView(APIView):
|
||||
# permission_classes = [permissions.IsAuthenticated]
|
||||
|
||||
Reference in New Issue
Block a user