From 0525f9559ba6f47fdb360c04c91453d3cddf3c63 Mon Sep 17 00:00:00 2001 From: Ryan Westfall Date: Sun, 26 Jul 2026 05:00:17 -0700 Subject: [PATCH] Add offline unit test suite for chat_backend (#12) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Closes #5 ## Summary - Replaces the three scattered test modules (`chat_backend/tests.py`, `services/tests.py`, `services/prompt_classifier/tests.py`) with a `chat_backend/tests/` package: **242 deterministic tests plus 6 opt-in live-Ollama checks**, up from 10 tests (3 of which were skipped and 4 of which were never even discovered). - The suite runs fully offline — no Ollama, Chroma, SMTP or network access. LangChain runnables are replaced by a small `FakeChain`, Chroma/embeddings are mocked, email uses Django's locmem backend, and blobs go through `DatabaseStorage`. - New `llm_be/test_runner.py` (wired via `TEST_RUNNER`) sets `SKIP_RAG_INIT=1` and an MD5 password hasher, so the suite cannot accidentally reach a model server and finishes in ~5s on SQLite (~17s on Postgres) instead of ~30s. ## Coverage | Area | File | |------|------| | `TimeInfoBase.save`, slugs, `get_duration`, `file_exists`, cascades | `test_models.py` | | `DatabaseStorage` save/open/exists/size/listdir/delete/times | `test_storage.py` | | JWT claim, prompt/user/feedback/document serializers | `test_serializers.py` | | auth + token, invite, feedback, company users, set-password, TOS | `test_views_users.py` | | conversation list/order/create/detail/soft-delete | `test_views_conversations.py` | | all four analytics endpoints, including empty-month behaviour | `test_views_analytics.py` | | workspace + document upload/list/detail, 404 and 400 paths | `test_views_documents.py` | | prompt classifier rules/parsing, moderation fail-safe, title cleanup | `test_services_classifiers.py` | | CSV/XLSX/DOCX/PDF analysis, plot generation, error payloads | `test_services_data_analysis.py` | | loader selection, filename sanitising, ingest, temp-file cleanup, search filters | `test_services_rag.py` | | history formatting and streaming | `test_services_llm.py` | | document re-index on create/delete, `SKIP_RAG_INIT` guard | `test_signals.py` | | consumer DB helpers, LangGraph nodes (moderation, classification, generation, search flags), websocket routes | `test_consumers.py` | Live checks (non-deterministic, need a model server): ```bash cd llm_be RUN_LIVE_OLLAMA_TESTS=1 uv run python manage.py test chat_backend.tests.test_live_ollama ``` ## Bugs the tests surfaced (fixed here) 1. **`ConversationDetailView.post` silently dropped every prompt.** `import datetime` shadowed `from datetime import datetime`, so `datetime.now()` raised `AttributeError` inside a bare `except` and the endpoint returned 200 without saving. Now uses `timezone.now()`. 2. **Prompt attachments never reached the LLM.** `get_conversation_file_async` (both consumers) did `sync_to_async(prompt.file.read)` — with `DatabaseStorage` the attribute access itself opens the blob, i.e. a DB query in async context, raising `SynchronousOnlyOperation` that was swallowed and returned `(None, None)`. The read now happens inside the thread. 3. **`DatabaseStorage._save` crashed on a str-backed `ContentFile`** (`TypeError: sequence item 0: expected a bytes-like object`); chunks are encoded when needed. 4. **`services/prompt_classifier/__init__.,py`** (note the comma) meant the directory was only an implicit namespace package, which is why its test module was never collected. Renamed, and its duplicate live-Ollama tests folded into `test_live_ollama.py`. Known-broken paths deliberately left untested and unchanged: `reset_password` / `ResetUserPassword` reference an unimported `requests` plus undefined locals, and `DocumentDetailView.get` references an undefined `workspaces` on its success path. Worth a follow-up ticket. ## Test plan - [x] `cd llm_be && uv run python manage.py test` → 248 tests, OK (6 skipped, all opt-in live) - [x] Same suite against Postgres 16 (`DATABASE_URL=postgres://…`) → OK, matching the containerized run in `deploy.yml` - [x] `uv run black` clean on all added files - [ ] Gitea Actions **Unit Tests** + **CI** green on this PRReviewed-on: https://git.aimloperations.com/ai_ml_operations/chat_backend/pulls/12 --- README.md | 17 +- llm_be/chat_backend/consumers.py | 4 +- llm_be/chat_backend/consumers_graph.py | 4 +- .../{__init__.,py => __init__.py} | 0 .../services/prompt_classifier/tests.py | 31 -- llm_be/chat_backend/services/tests.py | 32 -- llm_be/chat_backend/storage.py | 24 +- llm_be/chat_backend/tests.py | 194 --------- llm_be/chat_backend/tests/__init__.py | 6 + llm_be/chat_backend/tests/factories.py | 131 ++++++ llm_be/chat_backend/tests/fakes.py | 44 +++ llm_be/chat_backend/tests/test_consumers.py | 373 ++++++++++++++++++ llm_be/chat_backend/tests/test_live_ollama.py | 72 ++++ llm_be/chat_backend/tests/test_models.py | 194 +++++++++ llm_be/chat_backend/tests/test_serializers.py | 188 +++++++++ .../tests/test_services_classifiers.py | 221 +++++++++++ .../tests/test_services_data_analysis.py | 173 ++++++++ .../chat_backend/tests/test_services_llm.py | 64 +++ .../chat_backend/tests/test_services_rag.py | 279 +++++++++++++ llm_be/chat_backend/tests/test_signals.py | 54 +++ llm_be/chat_backend/tests/test_storage.py | 94 +++++ llm_be/chat_backend/tests/test_utils.py | 90 +++++ .../tests/test_views_analytics.py | 129 ++++++ .../tests/test_views_conversations.py | 141 +++++++ .../tests/test_views_documents.py | 142 +++++++ llm_be/chat_backend/tests/test_views_users.py | 354 +++++++++++++++++ llm_be/chat_backend/views.py | 10 +- llm_be/llm_be/settings.py | 6 +- llm_be/llm_be/test_runner.py | 20 + 29 files changed, 2813 insertions(+), 278 deletions(-) rename llm_be/chat_backend/services/prompt_classifier/{__init__.,py => __init__.py} (100%) delete mode 100644 llm_be/chat_backend/services/prompt_classifier/tests.py delete mode 100644 llm_be/chat_backend/services/tests.py delete mode 100644 llm_be/chat_backend/tests.py create mode 100644 llm_be/chat_backend/tests/__init__.py create mode 100644 llm_be/chat_backend/tests/factories.py create mode 100644 llm_be/chat_backend/tests/fakes.py create mode 100644 llm_be/chat_backend/tests/test_consumers.py create mode 100644 llm_be/chat_backend/tests/test_live_ollama.py create mode 100644 llm_be/chat_backend/tests/test_models.py create mode 100644 llm_be/chat_backend/tests/test_serializers.py create mode 100644 llm_be/chat_backend/tests/test_services_classifiers.py create mode 100644 llm_be/chat_backend/tests/test_services_data_analysis.py create mode 100644 llm_be/chat_backend/tests/test_services_llm.py create mode 100644 llm_be/chat_backend/tests/test_services_rag.py create mode 100644 llm_be/chat_backend/tests/test_signals.py create mode 100644 llm_be/chat_backend/tests/test_storage.py create mode 100644 llm_be/chat_backend/tests/test_utils.py create mode 100644 llm_be/chat_backend/tests/test_views_analytics.py create mode 100644 llm_be/chat_backend/tests/test_views_conversations.py create mode 100644 llm_be/chat_backend/tests/test_views_documents.py create mode 100644 llm_be/chat_backend/tests/test_views_users.py create mode 100644 llm_be/llm_be/test_runner.py diff --git a/README.md b/README.md index 9537ae9..817f351 100644 --- a/README.md +++ b/README.md @@ -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`). -Tests (skip live-Ollama classifier cases): +Tests: ```bash 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) diff --git a/llm_be/chat_backend/consumers.py b/llm_be/chat_backend/consumers.py index 312042f..89d6720 100644 --- a/llm_be/chat_backend/consumers.py +++ b/llm_be/chat_backend/consumers.py @@ -196,8 +196,8 @@ async def get_conversation_file_async(conversation_id): ).exclude(file='').order_by('created').afirst() if prompt_with_file and prompt_with_file.file: - # You must use sync_to_async to access the file's binary content - file_data = await sync_to_async(prompt_with_file.file.read)() + # Opening a DatabaseStorage file hits the DB, so read inside the thread. + file_data = await sync_to_async(lambda: prompt_with_file.file.read())() file_type = prompt_with_file.file_type return file_data, file_type except Exception as e: diff --git a/llm_be/chat_backend/consumers_graph.py b/llm_be/chat_backend/consumers_graph.py index fcf583a..21148a9 100644 --- a/llm_be/chat_backend/consumers_graph.py +++ b/llm_be/chat_backend/consumers_graph.py @@ -7,6 +7,7 @@ from typing import TypedDict, Annotated, List, Union, Dict, Any from django.utils import timezone from django.conf import settings from django.core.files.base import ContentFile +from asgiref.sync import sync_to_async from channels.generic.websocket import AsyncWebsocketConsumer from channels.db import database_sync_to_async from langchain_core.messages import HumanMessage, AIMessage, BaseMessage @@ -139,7 +140,8 @@ async def get_conversation_file_async(conversation_id): ).exclude(file='').order_by('created').afirst() if prompt_with_file and prompt_with_file.file: - file_data = await sync_to_async(prompt_with_file.file.read)() + # Opening a DatabaseStorage file hits the DB, so read inside the thread. + file_data = await sync_to_async(lambda: prompt_with_file.file.read())() file_type = prompt_with_file.file_type return file_data, file_type except Exception as e: diff --git a/llm_be/chat_backend/services/prompt_classifier/__init__.,py b/llm_be/chat_backend/services/prompt_classifier/__init__.py similarity index 100% rename from llm_be/chat_backend/services/prompt_classifier/__init__.,py rename to llm_be/chat_backend/services/prompt_classifier/__init__.py diff --git a/llm_be/chat_backend/services/prompt_classifier/tests.py b/llm_be/chat_backend/services/prompt_classifier/tests.py deleted file mode 100644 index 667ef4e..0000000 --- a/llm_be/chat_backend/services/prompt_classifier/tests.py +++ /dev/null @@ -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) \ No newline at end of file diff --git a/llm_be/chat_backend/services/tests.py b/llm_be/chat_backend/services/tests.py deleted file mode 100644 index 59a886e..0000000 --- a/llm_be/chat_backend/services/tests.py +++ /dev/null @@ -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) diff --git a/llm_be/chat_backend/storage.py b/llm_be/chat_backend/storage.py index ac915d6..773a1f2 100644 --- a/llm_be/chat_backend/storage.py +++ b/llm_be/chat_backend/storage.py @@ -27,13 +27,17 @@ class DatabaseStorage(Storage): def _save(self, name, content): name = self.get_available_name(name) if hasattr(content, "chunks"): - data = b"".join(chunk for chunk in content.chunks()) + chunks = content.chunks() else: - data = content.read() - if isinstance(data, str): - data = data.encode("utf-8") + chunks = [content.read()] + data = b"".join( + chunk.encode("utf-8") if isinstance(chunk, str) else chunk + for chunk in chunks + ) - content_type = getattr(content, "content_type", None) or mimetypes.guess_type(name)[0] + content_type = ( + getattr(content, "content_type", None) or mimetypes.guess_type(name)[0] + ) StoredFile = self._model() with transaction.atomic(): StoredFile.objects.update_or_create( @@ -56,8 +60,10 @@ class DatabaseStorage(Storage): prefix = path.rstrip("/") if prefix: prefix = f"{prefix}/" - names = self._model().objects.filter(name__startswith=prefix).values_list( - "name", flat=True + names = ( + self._model() + .objects.filter(name__startswith=prefix) + .values_list("name", flat=True) ) dirs: set[str] = set() files: list[str] = [] @@ -88,4 +94,6 @@ class DatabaseStorage(Storage): return self._model().objects.values_list("created", flat=True).get(name=name) def get_modified_time(self, name): - return self._model().objects.values_list("last_modified", flat=True).get(name=name) + return ( + self._model().objects.values_list("last_modified", flat=True).get(name=name) + ) diff --git a/llm_be/chat_backend/tests.py b/llm_be/chat_backend/tests.py deleted file mode 100644 index cdfbd13..0000000 --- a/llm_be/chat_backend/tests.py +++ /dev/null @@ -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) - - - - - diff --git a/llm_be/chat_backend/tests/__init__.py b/llm_be/chat_backend/tests/__init__.py new file mode 100644 index 0000000..1126b5d --- /dev/null +++ b/llm_be/chat_backend/tests/__init__.py @@ -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``. +""" diff --git a/llm_be/chat_backend/tests/factories.py b/llm_be/chat_backend/tests/factories.py new file mode 100644 index 0000000..350460a --- /dev/null +++ b/llm_be/chat_backend/tests/factories.py @@ -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)) diff --git a/llm_be/chat_backend/tests/fakes.py b/llm_be/chat_backend/tests/fakes.py new file mode 100644 index 0000000..f6b5c52 --- /dev/null +++ b/llm_be/chat_backend/tests/fakes.py @@ -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 diff --git a/llm_be/chat_backend/tests/test_consumers.py b/llm_be/chat_backend/tests/test_consumers.py new file mode 100644 index 0000000..5693cc8 --- /dev/null +++ b/llm_be/chat_backend/tests/test_consumers.py @@ -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/$"], + ) diff --git a/llm_be/chat_backend/tests/test_live_ollama.py b/llm_be/chat_backend/tests/test_live_ollama.py new file mode 100644 index 0000000..2743c19 --- /dev/null +++ b/llm_be/chat_backend/tests/test_live_ollama.py @@ -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) diff --git a/llm_be/chat_backend/tests/test_models.py b/llm_be/chat_backend/tests/test_models.py new file mode 100644 index 0000000..00ee191 --- /dev/null +++ b/llm_be/chat_backend/tests/test_models.py @@ -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) diff --git a/llm_be/chat_backend/tests/test_serializers.py b/llm_be/chat_backend/tests/test_serializers.py new file mode 100644 index 0000000..0ef05ac --- /dev/null +++ b/llm_be/chat_backend/tests/test_serializers.py @@ -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) diff --git a/llm_be/chat_backend/tests/test_services_classifiers.py b/llm_be/chat_backend/tests/test_services_classifiers.py new file mode 100644 index 0000000..70ea5e9 --- /dev/null +++ b/llm_be/chat_backend/tests/test_services_classifiers.py @@ -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", + ) diff --git a/llm_be/chat_backend/tests/test_services_data_analysis.py b/llm_be/chat_backend/tests/test_services_data_analysis.py new file mode 100644 index 0000000..2dde1ec --- /dev/null +++ b/llm_be/chat_backend/tests/test_services_data_analysis.py @@ -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"]) diff --git a/llm_be/chat_backend/tests/test_services_llm.py b/llm_be/chat_backend/tests/test_services_llm.py new file mode 100644 index 0000000..c415727 --- /dev/null +++ b/llm_be/chat_backend/tests/test_services_llm.py @@ -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}] + ) diff --git a/llm_be/chat_backend/tests/test_services_rag.py b/llm_be/chat_backend/tests/test_services_rag.py new file mode 100644 index 0000000..a14fbdb --- /dev/null +++ b/llm_be/chat_backend/tests/test_services_rag.py @@ -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), []) diff --git a/llm_be/chat_backend/tests/test_signals.py b/llm_be/chat_backend/tests/test_signals.py new file mode 100644 index 0000000..d9294e6 --- /dev/null +++ b/llm_be/chat_backend/tests/test_signals.py @@ -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) diff --git a/llm_be/chat_backend/tests/test_storage.py b/llm_be/chat_backend/tests/test_storage.py new file mode 100644 index 0000000..4187d00 --- /dev/null +++ b/llm_be/chat_backend/tests/test_storage.py @@ -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), {}) diff --git a/llm_be/chat_backend/tests/test_utils.py b/llm_be/chat_backend/tests/test_utils.py new file mode 100644 index 0000000..9bb91f9 --- /dev/null +++ b/llm_be/chat_backend/tests/test_utils.py @@ -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") diff --git a/llm_be/chat_backend/tests/test_views_analytics.py b/llm_be/chat_backend/tests/test_views_analytics.py new file mode 100644 index 0000000..aa3a985 --- /dev/null +++ b/llm_be/chat_backend/tests/test_views_analytics.py @@ -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) diff --git a/llm_be/chat_backend/tests/test_views_conversations.py b/llm_be/chat_backend/tests/test_views_conversations.py new file mode 100644 index 0000000..ef20fee --- /dev/null +++ b/llm_be/chat_backend/tests/test_views_conversations.py @@ -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) diff --git a/llm_be/chat_backend/tests/test_views_documents.py b/llm_be/chat_backend/tests/test_views_documents.py new file mode 100644 index 0000000..0f89e29 --- /dev/null +++ b/llm_be/chat_backend/tests/test_views_documents.py @@ -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") diff --git a/llm_be/chat_backend/tests/test_views_users.py b/llm_be/chat_backend/tests/test_views_users.py new file mode 100644 index 0000000..acdfa6a --- /dev/null +++ b/llm_be/chat_backend/tests/test_views_users.py @@ -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) diff --git a/llm_be/chat_backend/views.py b/llm_be/chat_backend/views.py index b6ae5ed..8c371df 100644 --- a/llm_be/chat_backend/views.py +++ b/llm_be/chat_backend/views.py @@ -30,7 +30,6 @@ from .models import ( ) from django.views.decorators.cache import never_cache from django.http import JsonResponse -from datetime import datetime from .client import LlamaClient from asgiref.sync import sync_to_async, async_to_sync from channels.generic.websocket import AsyncWebsocketConsumer @@ -73,19 +72,19 @@ from .services.data_analysis_service import AsyncDataAnalysisService from .ollama_config import ollama_llm_kwargs, ollama_model - from langchain_classic.chains import create_retrieval_chain from langchain_classic.chains.combine_documents import create_stuff_documents_chain from langchain_ollama import ChatOllama import logging -logger = logging.getLogger(__name__) +logger = logging.getLogger(__name__) CHANNEL_NAME: str = "llm_messages" MODEL_NAME: str = ollama_model() + # Create your views here. class CustomObtainTokenView(TokenObtainPairView): permission_classes = (permissions.AllowAny,) @@ -480,7 +479,7 @@ class ConversationDetailView(APIView): data={ "message": prompt, "user_created": is_user, - "created": datetime.now(), + "created": timezone.now(), } ) if serializer.is_valid(): @@ -725,9 +724,6 @@ llm = OllamaLLM(**ollama_llm_kwargs(model=MODEL_NAME)) # chain = prompt | llm.with_config({"run_name": "model"}) | output_parser.with_config({"run_name": "Assistant"}) - - - # Document Views class DocumentWorkspaceView(APIView): # permission_classes = [permissions.IsAuthenticated] diff --git a/llm_be/llm_be/settings.py b/llm_be/llm_be/settings.py index bd2b399..619c82c 100644 --- a/llm_be/llm_be/settings.py +++ b/llm_be/llm_be/settings.py @@ -120,7 +120,9 @@ CORS_ALLOWED_ORIGINS = env_list( # 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 -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", "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" TIME_ZONE = "UTC" USE_I18N = True diff --git a/llm_be/llm_be/test_runner.py b/llm_be/llm_be/test_runner.py new file mode 100644 index 0000000..2224620 --- /dev/null +++ b/llm_be/llm_be/test_runner.py @@ -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"]