Dockerize chat_backend + Ollama LAN + DB file storage (#6) (#7)
Unit Tests / test (push) Successful in 13s
Unit Tests / test (push) Successful in 13s
## Summary Implements [chat_backend#6](#6) Part A: - **uv** packaging (`pyproject.toml` + `uv.lock`), Docker/compose (dev + prod), entrypoint/validate-env, Gitea unit-test + auto-deploy workflows (mirror `scha`) - Env-driven Django settings (`DJANGO_*`, `DATABASE_URL`, CSRF/CORS) - **`OLLAMA_BASE_URL`** wired through all Ollama/LangChain clients (prod → `http://10.0.0.128:11434`) - **DatabaseStorage** — prompt/document file blobs in Postgres (`StoredFile`), not container FS; RAG materializes temp paths for loaders - ASGI via `gunicorn` + `UvicornWorker` (HTTP + WebSockets) Companion server-infra PR registers `app_catalog` / `host_apps` (port **8003**). ## Test plan - [ ] `uv sync && cd llm_be && SKIP_RAG_INIT=1 uv run python manage.py test` - [ ] `docker compose build && docker compose up` against bundled Postgres - [ ] Confirm Ollama calls use `OLLAMA_BASE_URL` (not hardcoded localhost) - [ ] Upload a document / prompt file → row in `chat_backend_storedfile`, no disk under `media/` - [ ] After server-infra merge + secret/Postgres/NPM: deploy via `deploy.sh --app chat_backend --env prod`Reviewed-on: #7
This commit was merged in pull request #7.
This commit is contained in:
@@ -1,18 +1,20 @@
|
||||
from abc import ABC, abstractmethod
|
||||
from langchain_ollama import OllamaLLM
|
||||
from langchain_core.output_parsers import StrOutputParser
|
||||
from django.conf import settings
|
||||
from chat_backend.ollama_config import ollama_llm_kwargs
|
||||
|
||||
|
||||
class BaseService(ABC):
|
||||
"""Abstract base class for LLM conversation services."""
|
||||
|
||||
def __init__(self, temperature=0.7):
|
||||
self.llm = OllamaLLM(
|
||||
model="llama3.2" if not settings.DEBUG else "gpt-oss:20b",
|
||||
temperature=0.7,
|
||||
top_k=50,
|
||||
top_p=0.9,
|
||||
repeat_penalty=1.1,
|
||||
num_ctx=4096,
|
||||
**ollama_llm_kwargs(
|
||||
temperature=temperature,
|
||||
top_k=50,
|
||||
top_p=0.9,
|
||||
repeat_penalty=1.1,
|
||||
num_ctx=4096,
|
||||
)
|
||||
)
|
||||
self.output_parser = StrOutputParser()
|
||||
self.output_parser = StrOutputParser()
|
||||
|
||||
@@ -11,6 +11,7 @@ from langchain_core.output_parsers import StrOutputParser
|
||||
import docx
|
||||
import pypdf
|
||||
from django.conf import settings
|
||||
from chat_backend.ollama_config import ollama_llm_kwargs
|
||||
|
||||
|
||||
class AsyncDataAnalysisService:
|
||||
@@ -19,9 +20,10 @@ class AsyncDataAnalysisService:
|
||||
def __init__(self):
|
||||
# A model with a large context window and strong analytical skills is best
|
||||
self.llm = OllamaLLM(
|
||||
model="llama3.2" if not settings.DEBUG else "gpt-oss:20b",
|
||||
temperature=0.3,
|
||||
num_ctx=8192,
|
||||
**ollama_llm_kwargs(
|
||||
temperature=0.3,
|
||||
num_ctx=8192,
|
||||
)
|
||||
)
|
||||
self.output_parser = StrOutputParser()
|
||||
self._setup_chain()
|
||||
|
||||
@@ -8,6 +8,7 @@ from langchain_core.prompts import ChatPromptTemplate
|
||||
from django.conf import settings
|
||||
|
||||
from chat_backend.models import Conversation, Prompt
|
||||
from chat_backend.ollama_config import ollama_llm_kwargs
|
||||
|
||||
|
||||
class LLMService(ABC):
|
||||
@@ -15,12 +16,13 @@ class LLMService(ABC):
|
||||
|
||||
def __init__(self):
|
||||
self.llm = OllamaLLM(
|
||||
model="llama3.2" if not settings.DEBUG else "gpt-oss:20b",
|
||||
temperature=0.7,
|
||||
top_k=50,
|
||||
top_p=0.9,
|
||||
repeat_penalty=1.1,
|
||||
num_ctx=4096,
|
||||
**ollama_llm_kwargs(
|
||||
temperature=0.7,
|
||||
top_k=50,
|
||||
top_p=0.9,
|
||||
repeat_penalty=1.1,
|
||||
num_ctx=4096,
|
||||
)
|
||||
)
|
||||
self.output_parser = StrOutputParser()
|
||||
|
||||
|
||||
@@ -3,6 +3,7 @@ from typing import Dict, Any
|
||||
from langchain_core.prompts import ChatPromptTemplate
|
||||
from langchain_ollama import OllamaLLM
|
||||
from chat_backend.services.base_service import BaseService
|
||||
from chat_backend.ollama_config import ollama_llm_kwargs
|
||||
|
||||
|
||||
class ModerationLabel(Enum):
|
||||
@@ -19,10 +20,11 @@ class ModerationClassifier(BaseService):
|
||||
def __init__(self):
|
||||
super().__init__(temperature=0.1)
|
||||
self.llm = OllamaLLM(
|
||||
model="llama3.2",
|
||||
temperature=0.1, # Very low for strict moderation
|
||||
top_k=10,
|
||||
num_ctx=2048,
|
||||
**ollama_llm_kwargs(
|
||||
temperature=0.1, # Very low for strict moderation
|
||||
top_k=10,
|
||||
num_ctx=2048,
|
||||
)
|
||||
)
|
||||
|
||||
self.moderation_prompt = ChatPromptTemplate.from_messages(
|
||||
|
||||
@@ -1,11 +1,12 @@
|
||||
import os
|
||||
import re
|
||||
import tempfile
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import List, Dict, Any, AsyncGenerator, Generator, Optional
|
||||
from channels.db import database_sync_to_async
|
||||
from langchain_community.embeddings import OllamaEmbeddings
|
||||
from langchain_ollama import OllamaEmbeddings
|
||||
from django.conf import settings
|
||||
|
||||
# from langchain_community.llms import Ollama
|
||||
from langchain_ollama import OllamaLLM
|
||||
from langchain_community.vectorstores import Chroma
|
||||
from langchain_core.documents import Document as LangDocument
|
||||
@@ -23,6 +24,7 @@ from django.core.files.uploadedfile import UploadedFile
|
||||
from chat_backend.models import Conversation, Prompt, DocumentWorkspace, Document
|
||||
from pathlib import Path
|
||||
from chat_backend.services.base_service import BaseService
|
||||
from chat_backend.ollama_config import ollama_embeddings_kwargs
|
||||
|
||||
|
||||
@database_sync_to_async
|
||||
@@ -45,7 +47,7 @@ class RAGService(BaseService):
|
||||
return cls._instance
|
||||
|
||||
def __init__(self):
|
||||
self.embedding_model = OllamaEmbeddings(model="llama3.2" if not settings.DEBUG else "gpt-oss:20b")
|
||||
self.embedding_model = OllamaEmbeddings(**ollama_embeddings_kwargs())
|
||||
super().__init__()
|
||||
self.text_splitter = RecursiveCharacterTextSplitter(
|
||||
chunk_size=1000, chunk_overlap=200
|
||||
@@ -63,7 +65,10 @@ class RAGService(BaseService):
|
||||
|
||||
def _initialize_vector_store(self) -> Chroma:
|
||||
"""Initialize and return the Chroma vector store."""
|
||||
persist_directory = f"./chroma_db/"
|
||||
persist_directory = getattr(
|
||||
settings, "CHROMA_PERSIST_DIRECTORY", "./chroma_db/"
|
||||
)
|
||||
os.makedirs(persist_directory, exist_ok=True)
|
||||
vector_store = Chroma(
|
||||
embedding_function=self.embedding_model, persist_directory=persist_directory
|
||||
)
|
||||
@@ -74,19 +79,54 @@ class RAGService(BaseService):
|
||||
self.vector_store.delete_collection()
|
||||
self.vector_store = self._initialize_vector_store()
|
||||
|
||||
def _materialize_file_field(self, file_field) -> str:
|
||||
"""
|
||||
Write DB-backed FileField bytes to a NamedTemporaryFile for loaders
|
||||
that require a filesystem path. Caller must os.unlink the path.
|
||||
"""
|
||||
suffix = Path(file_field.name).suffix
|
||||
tmp = tempfile.NamedTemporaryFile(delete=False, suffix=suffix)
|
||||
try:
|
||||
file_field.open("rb")
|
||||
try:
|
||||
while True:
|
||||
chunk = file_field.read(1024 * 1024)
|
||||
if not chunk:
|
||||
break
|
||||
tmp.write(chunk)
|
||||
finally:
|
||||
file_field.close()
|
||||
tmp.close()
|
||||
return tmp.name
|
||||
except Exception:
|
||||
tmp.close()
|
||||
if os.path.exists(tmp.name):
|
||||
os.unlink(tmp.name)
|
||||
raise
|
||||
|
||||
def _prepare_documents(self, documents: List[Document]) -> List[Document]:
|
||||
"""Process documents for ingestion into vector store."""
|
||||
docs = []
|
||||
|
||||
for doc in documents:
|
||||
print(f"Processing: {doc.file.name}")
|
||||
loader_class = self._get_file_loader(doc.file.name)
|
||||
loader = loader_class(doc.file)
|
||||
|
||||
chunks = self._load_and_split_documents(doc.file.path)
|
||||
if chunks:
|
||||
self.vector_store.add_documents(chunks)
|
||||
tmp_path = self._materialize_file_field(doc.file)
|
||||
try:
|
||||
chunks = self._load_and_split_documents(
|
||||
tmp_path,
|
||||
metadata={
|
||||
"source": doc.file.name,
|
||||
"workspace_id": doc.workspace_id,
|
||||
"document_id": doc.id,
|
||||
},
|
||||
)
|
||||
if chunks:
|
||||
self.vector_store.add_documents(chunks)
|
||||
finally:
|
||||
if os.path.exists(tmp_path):
|
||||
os.unlink(tmp_path)
|
||||
self.vector_store.persist()
|
||||
return docs
|
||||
|
||||
def ingest_documents(self, workspace: DocumentWorkspace | None = None) -> None:
|
||||
"""Ingest documents from a workspace into the vector store."""
|
||||
@@ -99,18 +139,6 @@ class RAGService(BaseService):
|
||||
print(f"Processing the documents : {documents}")
|
||||
self._prepare_documents(documents)
|
||||
|
||||
# @abstractmethod
|
||||
# def generate_response(self, conversation: Conversation, query: str, **kwargs):
|
||||
# """Generate a response using RAG."""
|
||||
# pass
|
||||
|
||||
# @abstractmethod
|
||||
# def search_documents(
|
||||
# self, query: str, workspace: Optional[DocumentWorkspace] = None, k: int = 4
|
||||
# ) -> List[Document]:
|
||||
# """Search relevant documents from the vector store."""
|
||||
# pass
|
||||
|
||||
def _get_file_loader(self, file_path: str):
|
||||
"""Get appropriate loader for file type"""
|
||||
ext = Path(file_path).suffix.lower()
|
||||
@@ -120,18 +148,6 @@ class RAGService(BaseService):
|
||||
"""Sanitize filename for safe storage"""
|
||||
return re.sub(r"[^\w\-_. ]", "_", filename)
|
||||
|
||||
def _save_uploaded_file(self, uploaded_file: UploadedFile, save_dir: str) -> str:
|
||||
"""Save uploaded file to disk"""
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
sanitized_name = self._sanitize_filename(uploaded_file.name)
|
||||
file_path = os.path.join(save_dir, sanitized_name)
|
||||
|
||||
with open(file_path, "wb+") as destination:
|
||||
for chunk in uploaded_file.chunks():
|
||||
destination.write(chunk)
|
||||
|
||||
return file_path
|
||||
|
||||
def _load_and_split_documents(
|
||||
self, file_path: str, metadata: dict = None
|
||||
) -> List[Document]:
|
||||
@@ -148,46 +164,47 @@ class RAGService(BaseService):
|
||||
|
||||
def add_files_to_store(
|
||||
self,
|
||||
file_tupls: List[UploadedFile], # (file_path, name,workspace_id)
|
||||
file_tupls: List, # (file_path_or_field, name, workspace_id)
|
||||
workspace_id: str,
|
||||
source: str = "upload",
|
||||
save_dir: str = "data/uploads",
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Process and add uploaded files to vector store
|
||||
Process and add files to vector store.
|
||||
|
||||
Args:
|
||||
files: List of Django UploadedFile objects
|
||||
workspace_id: ID of the workspace these belong to
|
||||
source: Source identifier for documents
|
||||
save_dir: Directory to save uploaded files
|
||||
|
||||
Returns:
|
||||
Dictionary with processing results
|
||||
file_tupls entries: (path_str | Django FileField, original_name, workspace_id)
|
||||
Paths may be temp files; FileFields are materialized from DB storage.
|
||||
"""
|
||||
results = {"total_added": 0, "failed_files": [], "processed_files": []}
|
||||
|
||||
for file_tuple in file_tupls:
|
||||
tmp_created = None
|
||||
try:
|
||||
# Save file to disk
|
||||
file_ref, original_name, ws_id = (
|
||||
file_tuple[0],
|
||||
file_tuple[1],
|
||||
file_tuple[2],
|
||||
)
|
||||
if isinstance(file_ref, str):
|
||||
file_path = file_ref
|
||||
else:
|
||||
tmp_created = self._materialize_file_field(file_ref)
|
||||
file_path = tmp_created
|
||||
|
||||
# Prepare metadata
|
||||
metadata = {
|
||||
"source": file_tuple[1],
|
||||
"workspace_id": file_tuple[2],
|
||||
"original_filename": file_tuple[1],
|
||||
"file_path": file_tuple[0],
|
||||
"source": original_name,
|
||||
"workspace_id": ws_id,
|
||||
"original_filename": original_name,
|
||||
"file_path": original_name,
|
||||
}
|
||||
|
||||
# Load and split documents
|
||||
docs = self._load_and_split_documents(file_path, metadata)
|
||||
|
||||
# Add to vector store
|
||||
if docs:
|
||||
self.vector_store.add_documents(docs)
|
||||
results["total_added"] += len(docs)
|
||||
results["processed_files"].append(
|
||||
{"filename": file_tuple[1], "document_count": len(docs)}
|
||||
{"filename": original_name, "document_count": len(docs)}
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
@@ -195,8 +212,10 @@ class RAGService(BaseService):
|
||||
{"filename": file_tuple[1], "error": str(e)}
|
||||
)
|
||||
continue
|
||||
finally:
|
||||
if tmp_created and os.path.exists(tmp_created):
|
||||
os.unlink(tmp_created)
|
||||
|
||||
# Persist changes
|
||||
self.vector_store.persist()
|
||||
return results
|
||||
|
||||
@@ -245,7 +264,6 @@ class SyncRAGService(RAGService):
|
||||
query = input_dict["query"]
|
||||
conversation = input_dict["conversation"]
|
||||
|
||||
# You could enhance this to consider historical context in retrieval
|
||||
relevant_docs = self.search_documents(query, conversation.workspace)
|
||||
if not relevant_docs:
|
||||
print("didn't find any relevant docs")
|
||||
@@ -260,10 +278,11 @@ class SyncRAGService(RAGService):
|
||||
filter_dict = {}
|
||||
if workspace:
|
||||
filter_dict["workspace_id"] = workspace.id
|
||||
search_kwargs = {"k": k, "filter": filter_dict if filter_dict else None}
|
||||
print(f"search_kwargs: {search_kwargs}")
|
||||
retriever = self.vector_store.as_retriever(
|
||||
search_type="similarity",
|
||||
search_kwargs={"k": k, "filter": filter_dict if filter_dict else None},
|
||||
search_kwargs=search_kwargs,
|
||||
)
|
||||
return retriever.get_relevant_documents(query)
|
||||
|
||||
@@ -299,7 +318,7 @@ class AsyncRAGService(RAGService):
|
||||
self.rag_chain = (
|
||||
{
|
||||
"context": self._retriever_with_history,
|
||||
"history": lambda x: x['recent_conversation'], #self._format_history(x["conversation"]),
|
||||
"history": lambda x: x["recent_conversation"],
|
||||
"question": lambda x: x["query"],
|
||||
}
|
||||
| self.prompt
|
||||
@@ -309,16 +328,12 @@ class AsyncRAGService(RAGService):
|
||||
|
||||
async def _format_history(self, conversation: Conversation) -> str:
|
||||
"""Format conversation history for the prompt."""
|
||||
# prompts = (
|
||||
# await Prompt.objects.filter(conversation=conversation)
|
||||
# .order_by("created_at")
|
||||
# .alist()
|
||||
# )
|
||||
# print(f"prompts that we are seeding with are: {prompts}")
|
||||
# return "\n".join(
|
||||
# f"{'User' if prompt.is_user else 'AI'}: {prompt.text}" for prompt in prompts
|
||||
# )
|
||||
return "\n".join([f"{"User" if prompt.type=="human" else "AI"}: {prompt.text()}" for prompt in conversation])
|
||||
return "\n".join(
|
||||
[
|
||||
f'{"User" if prompt.type == "human" else "AI"}: {prompt.text()}'
|
||||
for prompt in conversation
|
||||
]
|
||||
)
|
||||
|
||||
async def _retriever_with_history(self, input_dict: Dict[str, Any]) -> str:
|
||||
"""Retrieve documents considering conversation history."""
|
||||
@@ -327,7 +342,6 @@ class AsyncRAGService(RAGService):
|
||||
conversation = input_dict["conversation"]
|
||||
workspace = input_dict["workspace"]
|
||||
|
||||
# You could enhance this to consider historical context in retrieval
|
||||
docs = await self.search_documents(query, workspace)
|
||||
|
||||
if not docs:
|
||||
@@ -365,7 +379,7 @@ class AsyncRAGService(RAGService):
|
||||
"query": query,
|
||||
"conversation": conversation,
|
||||
"workspace": workspace,
|
||||
"recent_conversation": await self._format_history(conversation),
|
||||
"recent_conversation": await self._format_history(conversation),
|
||||
}
|
||||
|
||||
async for chunk in self.rag_chain.astream(chain_input):
|
||||
|
||||
@@ -1,251 +1,32 @@
|
||||
import os
|
||||
from unittest import TestCase, mock
|
||||
from unittest.mock import MagicMock, patch, AsyncMock
|
||||
from typing import List, Dict, Any
|
||||
import unittest
|
||||
from unittest import TestCase
|
||||
|
||||
from django.test import TestCase as DjangoTestCase
|
||||
|
||||
from chat_backend.services.rag_services import (
|
||||
RAGService,
|
||||
SyncRAGService,
|
||||
AsyncRAGService,
|
||||
from chat_backend.services.prompt_classifier.prompt_classifier import (
|
||||
PromptClassifier,
|
||||
PromptType,
|
||||
)
|
||||
from chat_backend.models import Conversation, Prompt, DocumentWorkspace, Document
|
||||
from chat_backend.services.prompt_classifier import PromptClassifier, PromptType
|
||||
from parameterized import parameterized
|
||||
|
||||
|
||||
# class TestRAGService(TestCase):
|
||||
# def setUp(self):
|
||||
# self.rag_service = RAGService()
|
||||
# self.rag_service.vector_store = MagicMock()
|
||||
# self.rag_service.embedding_model = MagicMock()
|
||||
# self.rag_service.text_splitter = MagicMock()
|
||||
|
||||
# def test_initialize_vector_store(self):
|
||||
# with patch("os.path.exists", return_value=False), patch(
|
||||
# "os.makedirs"
|
||||
# ) as mock_makedirs, patch(
|
||||
# "langchain_community.vectorstores.Chroma"
|
||||
# ) as mock_chroma:
|
||||
|
||||
# # Reset the vector store to test initialization
|
||||
# self.rag_service.vector_store = None
|
||||
# result = self.rag_service._initialize_vector_store()
|
||||
|
||||
# mock_makedirs.assert_called_once_with("chroma_db")
|
||||
# mock_chroma.assert_called_once_with(
|
||||
# embedding_function=self.rag_service.embedding_model,
|
||||
# persist_directory="chroma_db",
|
||||
# )
|
||||
# self.assertIsNotNone(result)
|
||||
|
||||
# def test_prepare_documents(self):
|
||||
# mock_doc1 = MagicMock(spec=Document)
|
||||
# mock_doc1.content = "Test content"
|
||||
# mock_doc1.source = "test_source"
|
||||
# mock_doc1.workspace = MagicMock()
|
||||
# mock_doc1.workspace.id = 1
|
||||
# mock_doc1.id = 1
|
||||
|
||||
# self.rag_service.text_splitter.split_text.return_value = ["chunk1", "chunk2"]
|
||||
|
||||
# result = self.rag_service._prepare_documents([mock_doc1])
|
||||
|
||||
# self.assertEqual(len(result), 2)
|
||||
# self.rag_service.text_splitter.split_text.assert_called_once_with(
|
||||
# "Test content"
|
||||
# )
|
||||
# self.assertEqual(result[0].page_content, "chunk1")
|
||||
# self.assertEqual(result[0].metadata["source"], "test_source")
|
||||
|
||||
# def test_ingest_documents(self):
|
||||
# mock_workspace = MagicMock()
|
||||
# mock_document = MagicMock()
|
||||
# mock_documents = [mock_document]
|
||||
|
||||
# with patch(
|
||||
# "services.rag_services.Document.objects.filter", return_value=mock_documents
|
||||
# ):
|
||||
# self.rag_service._prepare_documents = MagicMock(
|
||||
# return_value=["processed_doc"]
|
||||
# )
|
||||
|
||||
# self.rag_service.ingest_documents(mock_workspace)
|
||||
|
||||
# self.rag_service.vector_store.add_documents.assert_called_once_with(
|
||||
# ["processed_doc"]
|
||||
# )
|
||||
# self.rag_service.vector_store.persist.assert_called_once()
|
||||
|
||||
|
||||
# class TestSyncRAGService(DjangoTestCase):
|
||||
# def setUp(self):
|
||||
# self.sync_service = SyncRAGService()
|
||||
# self.sync_service.vector_store = MagicMock()
|
||||
# self.sync_service.llm = MagicMock()
|
||||
# self.sync_service.rag_chain = MagicMock()
|
||||
|
||||
# self.mock_conversation = MagicMock(spec=Conversation)
|
||||
# self.mock_conversation.workspace = MagicMock()
|
||||
|
||||
# self.mock_prompt1 = MagicMock(spec=Prompt)
|
||||
# self.mock_prompt1.is_user = True
|
||||
# self.mock_prompt1.text = "User question"
|
||||
# self.mock_prompt1.created_at = "2023-01-01"
|
||||
|
||||
# self.mock_prompt2 = MagicMock(spec=Prompt)
|
||||
# self.mock_prompt2.is_user = False
|
||||
# self.mock_prompt2.text = "AI response"
|
||||
# self.mock_prompt2.created_at = "2023-01-02"
|
||||
|
||||
# def test_format_history(self):
|
||||
# with patch("services.rag_services.Prompt.objects.filter") as mock_filter:
|
||||
# mock_filter.return_value.order_by.return_value = [
|
||||
# self.mock_prompt1,
|
||||
# self.mock_prompt2,
|
||||
# ]
|
||||
|
||||
# result = self.sync_service._format_history(self.mock_conversation)
|
||||
|
||||
# expected = "User: User question\nAI: AI response"
|
||||
# self.assertEqual(result, expected)
|
||||
# mock_filter.assert_called_once_with(conversation=self.mock_conversation)
|
||||
|
||||
# def test_retriever_with_history(self):
|
||||
# input_dict = {"query": "test query", "conversation": self.mock_conversation}
|
||||
|
||||
# self.sync_service.search_documents = MagicMock(return_value=["doc1", "doc2"])
|
||||
|
||||
# result = self.sync_service._retriever_with_history(input_dict)
|
||||
|
||||
# self.sync_service.search_documents.assert_called_once_with(
|
||||
# "test query", self.mock_conversation.workspace
|
||||
# )
|
||||
# self.assertEqual(result, ["doc1", "doc2"])
|
||||
|
||||
# def test_search_documents(self):
|
||||
# mock_retriever = MagicMock()
|
||||
# mock_retriever.get_relevant_documents.return_value = ["doc1", "doc2"]
|
||||
# self.sync_service.vector_store.as_retriever.return_value = mock_retriever
|
||||
|
||||
# result = self.sync_service.search_documents(
|
||||
# "test query", self.mock_conversation.workspace
|
||||
# )
|
||||
|
||||
# self.sync_service.vector_store.as_retriever.assert_called_once_with(
|
||||
# search_type="similarity",
|
||||
# search_kwargs={
|
||||
# "k": 4,
|
||||
# "filter": {"workspace_id": self.mock_conversation.workspace.id},
|
||||
# },
|
||||
# )
|
||||
# self.assertEqual(result, ["doc1", "doc2"])
|
||||
|
||||
# def test_generate_response(self):
|
||||
# chain_input = {"query": "test query", "conversation": self.mock_conversation}
|
||||
|
||||
# mock_stream = ["chunk1", "chunk2", "chunk3"]
|
||||
# self.sync_service.rag_chain.stream.return_value = mock_stream
|
||||
|
||||
# result = list(
|
||||
# self.sync_service.generate_response(self.mock_conversation, "test query")
|
||||
# )
|
||||
|
||||
# self.sync_service.rag_chain.stream.assert_called_once_with(chain_input)
|
||||
# self.assertEqual(result, mock_stream)
|
||||
|
||||
|
||||
# class TestAsyncRAGService(DjangoTestCase):
|
||||
# def setUp(self):
|
||||
# self.async_service = AsyncRAGService()
|
||||
# self.async_service.vector_store = MagicMock()
|
||||
# self.async_service.llm = MagicMock()
|
||||
# self.async_service.rag_chain = AsyncMock()
|
||||
|
||||
# self.mock_conversation = MagicMock(spec=Conversation)
|
||||
# self.mock_conversation.workspace = MagicMock()
|
||||
|
||||
# self.mock_prompt1 = MagicMock(spec=Prompt)
|
||||
# self.mock_prompt1.is_user = True
|
||||
# self.mock_prompt1.text = "User question"
|
||||
# self.mock_prompt1.created_at = "2023-01-01"
|
||||
|
||||
# self.mock_prompt2 = MagicMock(spec=Prompt)
|
||||
# self.mock_prompt2.is_user = False
|
||||
# self.mock_prompt2.text = "AI response"
|
||||
# self.mock_prompt2.created_at = "2023-01-02"
|
||||
|
||||
# async def test_format_history(self):
|
||||
# mock_manager = AsyncMock()
|
||||
# mock_manager.order_by.return_value.alist.return_value = [
|
||||
# self.mock_prompt1,
|
||||
# self.mock_prompt2,
|
||||
# ]
|
||||
|
||||
# with patch(
|
||||
# "services.rag_services.Prompt.objects.filter", return_value=mock_manager
|
||||
# ):
|
||||
# result = await self.async_service._format_history(self.mock_conversation)
|
||||
|
||||
# expected = "User: User question\nAI: AI response"
|
||||
# self.assertEqual(result, expected)
|
||||
# mock_manager.order_by.assert_called_once_with("created_at")
|
||||
|
||||
# async def test_retriever_with_history(self):
|
||||
# input_dict = {"query": "test query", "conversation": self.mock_conversation}
|
||||
|
||||
# self.async_service.search_documents = AsyncMock(return_value=["doc1", "doc2"])
|
||||
|
||||
# result = await self.async_service._retriever_with_history(input_dict)
|
||||
|
||||
# self.async_service.search_documents.assert_awaited_once_with(
|
||||
# "test query", self.mock_conversation.workspace
|
||||
# )
|
||||
# self.assertEqual(result, ["doc1", "doc2"])
|
||||
|
||||
# async def test_search_documents(self):
|
||||
# mock_retriever = AsyncMock()
|
||||
# mock_retriever.aget_relevant_documents.return_value = ["doc1", "doc2"]
|
||||
# self.async_service.vector_store.as_retriever.return_value = mock_retriever
|
||||
|
||||
# result = await self.async_service.search_documents(
|
||||
# "test query", self.mock_conversation.workspace
|
||||
# )
|
||||
|
||||
# self.async_service.vector_store.as_retriever.assert_called_once_with(
|
||||
# search_type="similarity",
|
||||
# search_kwargs={
|
||||
# "k": 4,
|
||||
# "filter": {"workspace_id": self.mock_conversation.workspace.id},
|
||||
# },
|
||||
# )
|
||||
# self.assertEqual(result, ["doc1", "doc2"])
|
||||
|
||||
# async def test_generate_response(self):
|
||||
# chain_input = {"query": "test query", "conversation": self.mock_conversation}
|
||||
|
||||
# mock_stream = ["chunk1", "chunk2", "chunk3"]
|
||||
# self.async_service.rag_chain.astream.return_value = mock_stream
|
||||
|
||||
# chunks = []
|
||||
# async for chunk in self.async_service.generate_response(
|
||||
# self.mock_conversation, "test query"
|
||||
# ):
|
||||
# chunks.append(chunk)
|
||||
|
||||
# self.async_service.rag_chain.astream.assert_awaited_once_with(chain_input)
|
||||
# self.assertEqual(chunks, mock_stream)
|
||||
|
||||
@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],
|
||||
])
|
||||
@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)
|
||||
self.assertEqual(result, expected_output)
|
||||
|
||||
@@ -3,6 +3,7 @@ from langchain_core.prompts import ChatPromptTemplate
|
||||
# from langchain_community.llms import Ollama
|
||||
from langchain_ollama import OllamaLLM
|
||||
from typing import Optional
|
||||
from chat_backend.ollama_config import ollama_llm_kwargs
|
||||
|
||||
|
||||
class TitleGenerator:
|
||||
@@ -12,10 +13,11 @@ class TitleGenerator:
|
||||
|
||||
def __init__(self):
|
||||
self.llm = OllamaLLM(
|
||||
model="llama3.2",
|
||||
temperature=0.5, # Slightly creative but not too random
|
||||
top_k=20,
|
||||
num_ctx=2048, # Shorter context needed for titles
|
||||
**ollama_llm_kwargs(
|
||||
temperature=0.5, # Slightly creative but not too random
|
||||
top_k=20,
|
||||
num_ctx=2048, # Shorter context needed for titles
|
||||
)
|
||||
)
|
||||
|
||||
self.title_prompt = ChatPromptTemplate.from_messages(
|
||||
|
||||
Reference in New Issue
Block a user