Add offline unit test suite for chat_backend (#5)
CI / test (pull_request) Successful in 10s
Unit Tests / test (pull_request) Successful in 9s

Replaces the three scattered test modules with a `chat_backend/tests/` package
covering models, DatabaseStorage, serializers, every REST endpoint, the LLM
services, signals and both websocket consumers — 242 deterministic tests that
run without Ollama, Chroma, SMTP or network access, plus 6 opt-in live-Ollama
checks behind RUN_LIVE_OLLAMA_TESTS.

A custom test runner sets SKIP_RAG_INIT and a cheap password hasher so the suite
finishes in seconds and can no longer reach a model server by accident.

Bugs the new tests surfaced, fixed here:
- views.ConversationDetailView.post silently dropped every prompt: `import
  datetime` shadowed `from datetime import datetime`, so `datetime.now()` raised
  inside a bare except.
- get_conversation_file_async in both consumers always returned None with
  DatabaseStorage, because `prompt.file.read` opens the blob (a DB query) during
  attribute access in async context.
- DatabaseStorage._save crashed on str-backed ContentFile.
- prompt_classifier/__init__.,py was never importable as a package init.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
2026-07-26 06:48:29 -05:00
co-authored by Cursor
parent 383c571137
commit bb4176590c
29 changed files with 2813 additions and 278 deletions
+15 -2
View File
@@ -44,11 +44,24 @@ uv run python manage.py runserver 0.0.0.0:8003
Without `DATABASE_URL` / `DB_HOST`, settings fall back to SQLite (`llm_be/db.sqlite3`). Without `DATABASE_URL` / `DB_HOST`, settings fall back to SQLite (`llm_be/db.sqlite3`).
Tests (skip live-Ollama classifier cases): Tests:
```bash ```bash
cd llm_be cd llm_be
SKIP_RAG_INIT=1 uv run python manage.py test uv run python manage.py test
```
The suite is offline by default — the custom test runner
(`llm_be/test_runner.py`) sets `SKIP_RAG_INIT=1` and a cheap password hasher, and
LLM chains are faked, so no Ollama, Chroma, SMTP or network access is needed.
Tests live in `llm_be/chat_backend/tests/` (models, storage, serializers, views,
services, signals, consumers).
Non-deterministic checks against a real model server are opt-in:
```bash
cd llm_be
RUN_LIVE_OLLAMA_TESTS=1 uv run python manage.py test chat_backend.tests.test_live_ollama
``` ```
### Docker (dev, bundled Postgres) ### Docker (dev, bundled Postgres)
+2 -2
View File
@@ -196,8 +196,8 @@ async def get_conversation_file_async(conversation_id):
).exclude(file='').order_by('created').afirst() ).exclude(file='').order_by('created').afirst()
if prompt_with_file and prompt_with_file.file: if prompt_with_file and prompt_with_file.file:
# You must use sync_to_async to access the file's binary content # Opening a DatabaseStorage file hits the DB, so read inside the thread.
file_data = await sync_to_async(prompt_with_file.file.read)() file_data = await sync_to_async(lambda: prompt_with_file.file.read())()
file_type = prompt_with_file.file_type file_type = prompt_with_file.file_type
return file_data, file_type return file_data, file_type
except Exception as e: 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.utils import timezone
from django.conf import settings from django.conf import settings
from django.core.files.base import ContentFile from django.core.files.base import ContentFile
from asgiref.sync import sync_to_async
from channels.generic.websocket import AsyncWebsocketConsumer from channels.generic.websocket import AsyncWebsocketConsumer
from channels.db import database_sync_to_async from channels.db import database_sync_to_async
from langchain_core.messages import HumanMessage, AIMessage, BaseMessage 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() ).exclude(file='').order_by('created').afirst()
if prompt_with_file and prompt_with_file.file: 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 file_type = prompt_with_file.file_type
return file_data, file_type return file_data, file_type
except Exception as e: 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): def _save(self, name, content):
name = self.get_available_name(name) name = self.get_available_name(name)
if hasattr(content, "chunks"): if hasattr(content, "chunks"):
data = b"".join(chunk for chunk in content.chunks()) chunks = content.chunks()
else: else:
data = content.read() chunks = [content.read()]
if isinstance(data, str): data = b"".join(
data = data.encode("utf-8") 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() StoredFile = self._model()
with transaction.atomic(): with transaction.atomic():
StoredFile.objects.update_or_create( StoredFile.objects.update_or_create(
@@ -56,8 +60,10 @@ class DatabaseStorage(Storage):
prefix = path.rstrip("/") prefix = path.rstrip("/")
if prefix: if prefix:
prefix = f"{prefix}/" prefix = f"{prefix}/"
names = self._model().objects.filter(name__startswith=prefix).values_list( names = (
"name", flat=True self._model()
.objects.filter(name__startswith=prefix)
.values_list("name", flat=True)
) )
dirs: set[str] = set() dirs: set[str] = set()
files: list[str] = [] files: list[str] = []
@@ -88,4 +94,6 @@ class DatabaseStorage(Storage):
return self._model().objects.values_list("created", flat=True).get(name=name) return self._model().objects.values_list("created", flat=True).get(name=name)
def get_modified_time(self, 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.views.decorators.cache import never_cache
from django.http import JsonResponse from django.http import JsonResponse
from datetime import datetime
from .client import LlamaClient from .client import LlamaClient
from asgiref.sync import sync_to_async, async_to_sync from asgiref.sync import sync_to_async, async_to_sync
from channels.generic.websocket import AsyncWebsocketConsumer 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 .ollama_config import ollama_llm_kwargs, ollama_model
from langchain_classic.chains import create_retrieval_chain from langchain_classic.chains import create_retrieval_chain
from langchain_classic.chains.combine_documents import create_stuff_documents_chain from langchain_classic.chains.combine_documents import create_stuff_documents_chain
from langchain_ollama import ChatOllama from langchain_ollama import ChatOllama
import logging import logging
logger = logging.getLogger(__name__)
logger = logging.getLogger(__name__)
CHANNEL_NAME: str = "llm_messages" CHANNEL_NAME: str = "llm_messages"
MODEL_NAME: str = ollama_model() MODEL_NAME: str = ollama_model()
# Create your views here. # Create your views here.
class CustomObtainTokenView(TokenObtainPairView): class CustomObtainTokenView(TokenObtainPairView):
permission_classes = (permissions.AllowAny,) permission_classes = (permissions.AllowAny,)
@@ -480,7 +479,7 @@ class ConversationDetailView(APIView):
data={ data={
"message": prompt, "message": prompt,
"user_created": is_user, "user_created": is_user,
"created": datetime.now(), "created": timezone.now(),
} }
) )
if serializer.is_valid(): 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"}) # chain = prompt | llm.with_config({"run_name": "model"}) | output_parser.with_config({"run_name": "Assistant"})
# Document Views # Document Views
class DocumentWorkspaceView(APIView): class DocumentWorkspaceView(APIView):
# permission_classes = [permissions.IsAuthenticated] # permission_classes = [permissions.IsAuthenticated]
+5 -1
View File
@@ -120,7 +120,9 @@ CORS_ALLOWED_ORIGINS = env_list(
# Ollama — GPU host on LAN for deployed envs; loopback for local Ollama. # Ollama — GPU host on LAN for deployed envs; loopback for local Ollama.
# Prod/beta control-node secret should set OLLAMA_BASE_URL=http://10.0.0.128:11434 # Prod/beta control-node secret should set OLLAMA_BASE_URL=http://10.0.0.128:11434
OLLAMA_BASE_URL = env("OLLAMA_BASE_URL", "http://127.0.0.1:11434") or "http://127.0.0.1:11434" OLLAMA_BASE_URL = (
env("OLLAMA_BASE_URL", "http://127.0.0.1:11434") or "http://127.0.0.1:11434"
)
OLLAMA_MODEL = env( OLLAMA_MODEL = env(
"OLLAMA_MODEL", "OLLAMA_MODEL",
"llama3.2" if not DEBUG else "gpt-oss:20b", "llama3.2" if not DEBUG else "gpt-oss:20b",
@@ -197,6 +199,8 @@ AUTH_PASSWORD_VALIDATORS = [
}, },
] ]
TEST_RUNNER = "llm_be.test_runner.ChatBackendTestRunner"
LANGUAGE_CODE = "en-us" LANGUAGE_CODE = "en-us"
TIME_ZONE = "UTC" TIME_ZONE = "UTC"
USE_I18N = True USE_I18N = True
+20
View File
@@ -0,0 +1,20 @@
"""Test runner tweaks shared by local runs and Gitea Actions."""
import os
from django.conf import settings
from django.test.runner import DiscoverRunner
class ChatBackendTestRunner(DiscoverRunner):
"""Keeps the suite fast and offline.
- Cheap password hasher: PBKDF2 otherwise dominates runtime.
- ``SKIP_RAG_INIT``: guards the Chroma/Ollama code paths that fire from
model signals, so the suite passes without a running Ollama.
"""
def setup_test_environment(self, **kwargs):
os.environ.setdefault("SKIP_RAG_INIT", "1")
super().setup_test_environment(**kwargs)
settings.PASSWORD_HASHERS = ["django.contrib.auth.hashers.MD5PasswordHasher"]