Add offline unit test suite for chat_backend (#12)
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:
2026-07-26 05:00:17 -07:00
parent 383c571137
commit 0525f9559b
29 changed files with 2813 additions and 278 deletions
+2 -2
View File
@@ -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:
+3 -1
View File
@@ -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)
-32
View File
@@ -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)
+16 -8
View File
@@ -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)
)
-194
View File
@@ -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)
+6
View File
@@ -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``.
"""
+131
View File
@@ -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))
+44
View File
@@ -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
+373
View File
@@ -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)
+194
View File
@@ -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), [])
+54
View File
@@ -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)
+94
View File
@@ -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), {})
+90
View File
@@ -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)
+3 -7
View File
@@ -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]