Always-on grounded retrieval + role-scoped Ollama models (#62 Phases 1–3) #65

Merged
westfarn merged 2 commits from feature/accuracy-grounded-retrieval-62 into master 2026-08-02 11:46:02 -07:00
32 changed files with 1734 additions and 187 deletions
+16 -2
View File
@@ -12,8 +12,22 @@ DATABASE_URL=postgres://chat_backend:chat_backend@db:5432/chat_backend
# Ollama — local loopback when Ollama runs on this machine; LAN IP for GPU host.
OLLAMA_BASE_URL=http://127.0.0.1:11434
# OLLAMA_MODEL=llama3.2
# OLLAMA_EMBED_MODEL=llama3.2
# Legacy fallback (used when role-specific vars unset). Prefer the role vars.
# OLLAMA_MODEL=gpt-oss:20b
# OLLAMA_MODEL_THINKING=gpt-oss:20b
# OLLAMA_MODEL_FAST=gemma4:latest
# OLLAMA_MODEL_UTILITY=llama3.2
# OLLAMA_EMBED_MODEL=nomic-embed-text
# OLLAMA_NUM_CTX_THINKING=16384
# OLLAMA_NUM_CTX_FAST=8192
# OLLAMA_NUM_CTX_UTILITY=4096
# Web search (#62) — SearxNG primary, DDGS failover. See README "SearxNG".
ALLOW_INTERNET_ACCESS=true
SEARCH_PROVIDER=searxng
SEARCH_FAILOVER_PROVIDER=ddgs
SEARXNG_BASE_URL=http://127.0.0.1:8088
# SEARXNG_TIMEOUT_SECONDS=8
# Email (SMTP2GO) — optional for local
EMAIL_HOST=mail.smtp2go.com
+16 -2
View File
@@ -32,8 +32,22 @@ WEB_PORT=8003
# Ollama on GPU host (ai-server-4080). Firewall must allow 10.0.0.0/24 → :11434.
OLLAMA_BASE_URL=http://10.0.0.128:11434
OLLAMA_MODEL=llama3.2
OLLAMA_EMBED_MODEL=llama3.2
# Role-scoped models (#62). After changing OLLAMA_EMBED_MODEL, run:
# python manage.py reindex_embeddings
OLLAMA_MODEL=gpt-oss:20b
OLLAMA_MODEL_THINKING=gpt-oss:20b
OLLAMA_MODEL_FAST=gemma4:latest
OLLAMA_MODEL_UTILITY=llama3.2
OLLAMA_EMBED_MODEL=nomic-embed-text
OLLAMA_NUM_CTX_THINKING=16384
OLLAMA_NUM_CTX_FAST=8192
# Web search (#62) — self-hosted SearxNG (recommended). DDGS is automatic failover.
ALLOW_INTERNET_ACCESS=true
SEARCH_PROVIDER=searxng
SEARCH_FAILOVER_PROVIDER=ddgs
# Point at the SearxNG container/service on the LAN (see README "SearxNG").
SEARXNG_BASE_URL=http://10.0.0.128:8088
# Email (SMTP2GO)
EMAIL_HOST=mail.smtp2go.com
+76 -1
View File
@@ -87,7 +87,14 @@ with `COMPOSE_DATABASE_URL` if needed.
| `DATABASE_URL` | SQLite fallback | yes | Shared Postgres in prod |
| `WEB_PORT` | n/a (compose maps 8003) | `8003` | Host port for prod compose |
| `OLLAMA_BASE_URL` | `http://127.0.0.1:11434` | yes | GPU host in prod: `http://10.0.0.128:11434` |
| `OLLAMA_MODEL` / `OLLAMA_EMBED_MODEL` | from `DEBUG` | optional | Override model names |
| `OLLAMA_MODEL` | `gpt-oss:20b` | optional | Legacy fallback for THINKING |
| `OLLAMA_MODEL_THINKING` / `_FAST` / `_UTILITY` | see defaults | optional | Role-scoped chat models (#62) |
| `OLLAMA_EMBED_MODEL` | `nomic-embed-text` | optional | Never falls back to a chat model |
| `OLLAMA_NUM_CTX_THINKING` / `_FAST` | `16384` / `8192` | optional | Context window per role |
| `ALLOW_INTERNET_ACCESS` | `true` | optional | Gate for live web retrieval |
| `SEARCH_PROVIDER` | `searxng` | optional | Primary search provider (#62) |
| `SEARCH_FAILOVER_PROVIDER` | `ddgs` | optional | Automatic failover |
| `SEARXNG_BASE_URL` | `http://127.0.0.1:8088` | yes if using SearxNG | Self-hosted SearxNG JSON API |
| `EMAIL_HOST_*` | empty | yes (prod/beta) | SMTP2GO |
| `CAPTCHA_SECRET_KEY` | empty | recommended | |
| `ENABLE_ACCOUNT_REGISTRATION` | `false` | optional | Self-serve sign-up; keep false until ready |
@@ -136,6 +143,74 @@ All clients (`ollama.Client`, `OllamaLLM`, `OllamaEmbeddings`, `ChatOllama`) use
Firewall / Ollama listen on ai-server-4080 must allow `10.0.0.0/24``:11434`.
### Role-scoped models (#62)
| Role | Setting | Default | Used for |
|------|---------|---------|----------|
| THINKING | `OLLAMA_MODEL_THINKING` | `gpt-oss:20b` | Default chat / grounded answers |
| FAST | `OLLAMA_MODEL_FAST` | `gemma4:latest` | FE `modelName=FAST` (smaller/faster — still grounded) |
| UTILITY | `OLLAMA_MODEL_UTILITY` | `llama3.2` | Classify / moderate / title / grounding decision |
| EMBED | `OLLAMA_EMBED_MODEL` | `nomic-embed-text` | Chroma embeddings |
After changing `OLLAMA_EMBED_MODEL`, rebuild the vector store (dimension change):
```bash
SKIP_RAG_INIT=1 uv run python manage.py reindex_embeddings
```
### SearxNG (web search)
Grounded chat uses a self-hosted **SearxNG** instance as the primary search
provider (`SEARCH_PROVIDER=searxng`), with DuckDuckGo (`ddgs`) as automatic
failover. Point `SEARXNG_BASE_URL` at the JSON API (no trailing path).
**Recommended: run SearxNG on the GPU/infra host next to Ollama**
(`10.0.0.128`), reachable from the chat_backend containers on the LAN.
Minimal compose snippet (add to `server-infra` or run on ai-server-4080):
```yaml
services:
searxng:
image: searxng/searxng:latest
restart: unless-stopped
ports:
- "8080:8080"
volumes:
- ./searxng:/etc/searxng:rw
environment:
- SEARXNG_BASE_URL=http://10.0.0.128:8088/
```
In `searxng/settings.yml` (created on first start), enable the JSON format:
```yaml
search:
formats:
- html
- json
```
Then set in `chat_backend_prod.env` / `chat_backend_beta.env`:
```text
ALLOW_INTERNET_ACCESS=true
SEARCH_PROVIDER=searxng
SEARCH_FAILOVER_PROVIDER=ddgs
SEARXNG_BASE_URL=http://10.0.0.128:8088
```
Firewall: allow `10.0.0.0/24``:8080` on the SearxNG host (same pattern as
Ollama `:11434`). Verify from a backend container:
```bash
curl -sG 'http://10.0.0.128:8088/search' --data-urlencode 'q=test' -d 'format=json' | head
```
If SearxNG is down, chat still works for non-factual turns; factual turns that
require retrieval return an explicit "couldn't reach live sources" message
instead of hallucinating from parametric memory.
## File storage
Prompt attachments and workspace documents use **`DatabaseStorage`**
+12 -1
View File
@@ -3,6 +3,9 @@ from django.conf import settings
from django.db import OperationalError, ProgrammingError
import os
import sys
import logging
logger = logging.getLogger(__name__)
class ChatBackendConfig(AppConfig):
@@ -20,6 +23,7 @@ class ChatBackendConfig(AppConfig):
"test",
"shell",
"check",
"reindex_embeddings",
}
if any(cmd in sys.argv for cmd in management_cmds):
return
@@ -29,7 +33,10 @@ class ChatBackendConfig(AppConfig):
FORCE_RELOAD = False
try:
from .services.rag_services import AsyncRAGService
from .services.rag_services import (
AsyncRAGService,
EmbeddingDimensionMismatch,
)
from chat_backend.models import Document
if Document.objects.exists():
@@ -41,6 +48,10 @@ class ChatBackendConfig(AppConfig):
if FORCE_RELOAD:
print("Force Reload ChromaDB with existing documents...")
rag_service.clear_vector_store()
except EmbeddingDimensionMismatch as exc:
# Loud failure — do not silently serve with the wrong embed model.
logger.error("RAG embedding dimension mismatch: %s", exc)
raise
except (OperationalError, ProgrammingError):
# Database tables might not exist yet during migration
pass
+49 -30
View File
@@ -13,15 +13,18 @@ from asgiref.sync import sync_to_async, async_to_sync
from langchain_core.messages import HumanMessage, AIMessage
from langchain_community.vectorstores import Chroma
from langchain_ollama import OllamaEmbeddings
from langchain_community.tools import DuckDuckGoSearchRun
from chat_backend.ollama_config import ollama_embeddings_kwargs
from chat_backend.ollama_config import (
ollama_embeddings_kwargs,
ollama_model_for_role,
resolve_chat_role,
)
from django.conf import settings as django_settings
from langchain_core.runnables import RunnableLambda, RunnableBranch, RunnablePassthrough
from langchain_core.tracers.context import collect_runs
from .models import Conversation, Prompt, PromptMetric, DocumentWorkspace, Document, CustomUser
from .serializers import PromptSerializer
from .services.llm_service import AsyncLLMService
from .services.llm_service import AsyncLLMService, build_chat_service
from .services.rag_services import AsyncRAGService
from .services.chat_tenant_scope import (
ChatTenantScopeError,
@@ -35,6 +38,7 @@ from .services.title_generator import title_generator
from .services.moderation_classifier import moderation_classifier, ModerationLabel
from .services.prompt_classifier.prompt_classifier import PromptClassifier, PromptType
from .services.data_analysis_service import AsyncDataAnalysisService
from .services.grounded_chat import citations_frame, prepare_grounded_chat
from .utils import (
TokenUsageCollector,
aiter_text_chunks,
@@ -52,7 +56,6 @@ from finance.services.quotas import (
logger = logging.getLogger(__name__)
CHANNEL_NAME: str = "llm_messages"
MODEL_NAME: str = "llama3.2"
PROMPT_CLASSIFIER = PromptClassifier()
@@ -175,7 +178,7 @@ def get_messages(conversation_id, prompt, file_string: str = None, file_type: st
@database_sync_to_async
def save_generated_message(conversation_id, message):
def save_generated_message(conversation_id, message, citations=None):
conversation = Conversation.objects.get(id=conversation_id)
# add the prompt to the conversation
@@ -189,7 +192,12 @@ def save_generated_message(conversation_id, message):
if serializer.is_valid():
prompt_instance = serializer.save()
prompt_instance.conversation_id = conversation.id
if citations:
prompt_instance.citations = citations
prompt_instance = serializer.save()
# Ensure citations survive even if serializer omits write.
if citations is not None:
Prompt.objects.filter(pk=prompt_instance.pk).update(citations=citations)
else:
print(serializer.errors)
@@ -478,24 +486,6 @@ class ChatConsumerAgain(AsyncWebsocketConsumer):
}
return {"type": "text", "content": "Image Generation is not supported at this time, but it will be soon."}
if prompt_type == PromptType.SEARCH:
# Check modelName first - if FAST, we skip search regardless of settings
if input_dict.get("model_name") == "FAST":
pass # Skip search
elif getattr(settings, "ALLOW_INTERNET_ACCESS", False):
try:
search = DuckDuckGoSearchRun()
search_results = search.run(input_dict["message"])
messages.append(HumanMessage(content=f"Search Results: {search_results}"))
except Exception as e:
logger.error(f"Search failed: {e}")
# If search fails, we proceed without it, essentially falling back to general chat
pass
else:
# If search is disabled, we could notify the user, but for now we'll just proceed
# potentially adding a system message or just letting the LLM handle it with its training data
pass
if prompt_type == PromptType.RAG:
try:
await enforce_feature_gate(chat_user, "rag")
@@ -520,9 +510,21 @@ class ChatConsumerAgain(AsyncWebsocketConsumer):
return {"type": "text", "content": "Please upload a file to perform data analysis."}
return service.generate_response(prompt_instance.message, decoded_file, file_type)
else: # GENERAL_CHAT or others
service = AsyncLLMService()
return service.generate_response(messages, prompt_instance.message, conversation_id)
else:
# GENERAL_CHAT / SEARCH / UNKNOWN — always-on grounding (#62).
# FAST selects a smaller model; it no longer skips search.
grounded = await prepare_grounded_chat(
message=input_dict["message"],
messages=messages,
model_name=input_dict.get("model_name"),
conversation_id=conversation_id,
)
if grounded.error:
return grounded.error
# Stash citations/model on the input for the caller.
input_dict["_citations"] = grounded.citations
input_dict["_resolved_model"] = grounded.model_name
return grounded.generator
# --- Execution ---
@@ -546,12 +548,13 @@ class ChatConsumerAgain(AsyncWebsocketConsumer):
# messages = messages[:-1] + [HumanMessage(content=altered_message)]
# I'll add it to the input_dict if needed.
resolved_model = ollama_model_for_role(resolve_chat_role(model))
prompt_metric = await create_prompt_metric(
prompt_instance.id,
prompt_instance.message,
True if file else False,
file_type,
MODEL_NAME,
resolved_model,
conversation_id,
)
@@ -562,7 +565,9 @@ class ChatConsumerAgain(AsyncWebsocketConsumer):
"file_type": file_type,
"messages": messages,
"prompt_instance": prompt_instance,
"model_name": model
"model_name": model,
"_citations": [],
"_resolved_model": resolved_model,
}
# Run the pipeline steps manually to handle the async generator return type of generate_response_step
@@ -602,8 +607,22 @@ class ChatConsumerAgain(AsyncWebsocketConsumer):
tokens_in, tokens_out = usage.pair
await self.send("END_OF_THE_STREAM_ENDER_GAME_42")
await save_generated_message(conversation_id, full_response)
citations = step2.get("_citations") or []
if citations:
await self.send_json_message(json.dumps(citations_frame(citations)))
# Prefer model actually used by the grounded path when present.
final_model = step2.get("_resolved_model") or resolved_model
if final_model and final_model != prompt_metric.model_name:
prompt_metric.model_name = final_model
await database_sync_to_async(prompt_metric.save)(
update_fields=["model_name"]
)
await save_generated_message(
conversation_id, full_response, citations=citations
)
await finish_prompt_metric(
prompt_metric,
len(full_response),
+47 -26
View File
@@ -11,7 +11,6 @@ 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
from langchain_community.tools import DuckDuckGoSearchRun
from langgraph.graph import StateGraph, END
from .models import Conversation, Prompt, PromptMetric, DocumentWorkspace, CustomUser
@@ -30,6 +29,8 @@ from .services.title_generator import title_generator
from .services.moderation_classifier import moderation_classifier, ModerationLabel
from .services.prompt_classifier.prompt_classifier import PromptClassifier, PromptType
from .services.data_analysis_service import AsyncDataAnalysisService
from .services.grounded_chat import citations_frame, prepare_grounded_chat
from chat_backend.ollama_config import ollama_model_for_role, resolve_chat_role
from .utils import (
TokenUsageCollector,
aiter_text_chunks,
@@ -43,7 +44,6 @@ from finance.services.quotas import FeatureNotAllowed, QuotaExceeded, check_gene
logger = logging.getLogger(__name__)
CHANNEL_NAME: str = "llm_messages"
MODEL_NAME: str = "llama3.2"
PROMPT_CLASSIFIER = PromptClassifier()
# --- Database Helpers (Reused) ---
@@ -146,7 +146,7 @@ def get_messages(conversation_id, prompt, file_string: str = None, file_type: st
return transformed_messages, prompt_instance
@database_sync_to_async
def save_generated_message(conversation_id, message):
def save_generated_message(conversation_id, message, citations=None):
conversation = Conversation.objects.get(id=conversation_id)
serializer = PromptSerializer(
data={
@@ -159,6 +159,8 @@ def save_generated_message(conversation_id, message):
prompt_instance = serializer.save()
prompt_instance.conversation_id = conversation.id
prompt_instance.save()
if citations is not None:
Prompt.objects.filter(pk=prompt_instance.pk).update(citations=citations)
else:
print(serializer.errors)
@@ -221,6 +223,8 @@ class ChatState(TypedDict):
error: Union[str, None]
model_name: str
chat_user: Any
citations: List[Dict[str, Any]]
resolved_model: str
# --- LangGraph Nodes ---
@@ -277,22 +281,7 @@ async def generation_node(state: ChatState) -> ChatState:
}
return {"response_generator": {"type": "text", "content": "Image Generation is not supported at this time, but it will be soon."}}
# Feature Flag: Internet Access
if prompt_type == PromptType.SEARCH:
# Check modelName first - if FAST, we skip search regardless of settings
if state.get("model_name") == "FAST":
pass
elif getattr(settings, "ALLOW_INTERNET_ACCESS", False):
try:
search = DuckDuckGoSearchRun()
search_results = search.run(state["message"])
messages.append(HumanMessage(content=f"Search Results: {search_results}"))
except Exception as e:
logger.error(f"Search failed: {e}")
pass
else:
pass
# Feature Flag: Internet Access / always-on grounding handled below for chat.
if prompt_type == PromptType.RAG:
chat_user = state.get("chat_user")
if chat_user is not None:
@@ -318,10 +307,26 @@ async def generation_node(state: ChatState) -> ChatState:
generator = service.generate_response(prompt_instance.message, decoded_file, file_type)
return {"response_generator": generator}
else: # GENERAL_CHAT or others
service = AsyncLLMService()
generator = service.generate_response(messages, prompt_instance.message, conversation_id)
return {"response_generator": generator}
else:
# GENERAL_CHAT / SEARCH / UNKNOWN — always-on grounding (#62).
# FAST selects a smaller model; it no longer skips search.
grounded = await prepare_grounded_chat(
message=state["message"],
messages=messages,
model_name=state.get("model_name"),
conversation_id=conversation_id,
)
if grounded.error:
return {
"response_generator": grounded.error,
"citations": [],
"resolved_model": grounded.model_name or "",
}
return {
"response_generator": grounded.generator,
"citations": grounded.citations,
"resolved_model": grounded.model_name,
}
# --- LangGraph Definition ---
@@ -475,12 +480,13 @@ class ChatConsumerGraph(AsyncWebsocketConsumer):
if not decoded_file:
decoded_file, file_type = await get_conversation_file_async(conversation_id)
resolved_model = ollama_model_for_role(resolve_chat_role(model))
prompt_metric = await create_prompt_metric(
prompt_instance.id,
prompt_instance.message,
True if file else False,
file_type,
MODEL_NAME,
resolved_model,
conversation_id,
)
@@ -498,6 +504,8 @@ class ChatConsumerGraph(AsyncWebsocketConsumer):
"error": None,
"model_name": model,
"chat_user": chat_user,
"citations": [],
"resolved_model": resolved_model,
}
print("Initial State: ", initial_state)
@@ -535,8 +543,21 @@ class ChatConsumerGraph(AsyncWebsocketConsumer):
tokens_in, tokens_out = usage.pair
await self.send("END_OF_THE_STREAM_ENDER_GAME_42")
await save_generated_message(conversation_id, full_response)
citations = final_state.get("citations") or []
if citations:
await self.send_json_message(json.dumps(citations_frame(citations)))
final_model = final_state.get("resolved_model") or resolved_model
if final_model and final_model != prompt_metric.model_name:
prompt_metric.model_name = final_model
await database_sync_to_async(prompt_metric.save)(
update_fields=["model_name"]
)
await save_generated_message(
conversation_id, full_response, citations=citations
)
await finish_prompt_metric(
prompt_metric,
len(full_response),
@@ -0,0 +1,53 @@
"""Rebuild the Chroma collection with the configured embedding model (#62).
Usage:
python manage.py reindex_embeddings
python manage.py reindex_embeddings --dry-run
"""
from __future__ import annotations
from django.core.management.base import BaseCommand
from chat_backend.models import Document
from chat_backend.ollama_config import ollama_embed_model
from chat_backend.services.rag_services import AsyncRAGService
class Command(BaseCommand):
help = (
"Drop and rebuild the Chroma vector store using OLLAMA_EMBED_MODEL, "
"re-ingesting every Document while preserving workspace/document metadata."
)
def add_arguments(self, parser):
parser.add_argument(
"--dry-run",
action="store_true",
help="Print what would be re-ingested without mutating Chroma.",
)
def handle(self, *args, **options):
dry_run = options.get("dry_run")
embed_model = ollama_embed_model()
total = Document.objects.count()
active = Document.objects.filter(active=True).count()
self.stdout.write(
f"OLLAMA_EMBED_MODEL={embed_model} documents={total} "
f"(active={active})"
)
if dry_run:
self.stdout.write(self.style.WARNING("Dry run — no changes made."))
return
# Reset singleton so a fresh store is built under the current embed model.
AsyncRAGService._instance = None
rag = AsyncRAGService()
self.stdout.write("Clearing Chroma collection…")
rag.clear_vector_store()
self.stdout.write("Re-ingesting documents…")
rag.ingest_documents()
count = rag.vector_store._collection.count()
self.stdout.write(
self.style.SUCCESS(f"Reindex complete. Vector chunks now: {count}")
)
@@ -0,0 +1,18 @@
# Generated by Django 6.0 on 2026-08-02 14:03
from django.db import migrations, models
class Migration(migrations.Migration):
dependencies = [
('chat_backend', '0030_drive_sync_progress_counters'),
]
operations = [
migrations.AddField(
model_name='prompt',
name='citations',
field=models.JSONField(blank=True, default=list, help_text='Structured source citations for grounded answers (#62). List of {index, title, url, published_at}.'),
),
]
+8
View File
@@ -310,6 +310,14 @@ class Prompt(TimeInfoBase):
null=True,
help_text="file type of the file for the prompt",
)
citations = models.JSONField(
default=list,
blank=True,
help_text=(
"Structured source citations for grounded answers (#62). "
"List of {index, title, url, published_at}."
),
)
def get_conversation_title(self):
if self.conversation:
+77 -5
View File
@@ -1,25 +1,97 @@
"""Shared Ollama client helpers — always use settings.OLLAMA_BASE_URL."""
"""Shared Ollama client helpers — always use settings.OLLAMA_BASE_URL.
Role-scoped models (#62): THINKING / FAST / UTILITY / EMBED each resolve
independently, with ``OLLAMA_MODEL`` kept as a fallback so existing deploys
keep working until they set the role-specific vars.
"""
from __future__ import annotations
from django.conf import settings
ROLE_THINKING = "thinking"
ROLE_FAST = "fast"
ROLE_UTILITY = "utility"
ROLE_EMBED = "embed"
_VALID_ROLES = {ROLE_THINKING, ROLE_FAST, ROLE_UTILITY, ROLE_EMBED}
def ollama_base_url() -> str:
return getattr(settings, "OLLAMA_BASE_URL", "http://127.0.0.1:11434")
def ollama_model(default: str | None = None) -> str:
"""Legacy single-model accessor. Prefer :func:`ollama_model_for_role`."""
if default:
return default
return getattr(settings, "OLLAMA_MODEL", "llama3.2")
return getattr(settings, "OLLAMA_MODEL", "gpt-oss:20b")
def ollama_model_for_role(role: str) -> str:
"""Resolve the model name for a generation role.
Lookup order: role-specific setting → ``OLLAMA_MODEL`` fallback →
hard-coded role default (never falls back from embed → chat model).
"""
role = (role or ROLE_THINKING).lower()
if role not in _VALID_ROLES:
raise ValueError(f"Unknown Ollama role: {role!r}")
role_setting = {
ROLE_THINKING: "OLLAMA_MODEL_THINKING",
ROLE_FAST: "OLLAMA_MODEL_FAST",
ROLE_UTILITY: "OLLAMA_MODEL_UTILITY",
ROLE_EMBED: "OLLAMA_EMBED_MODEL",
}[role]
role_default = {
ROLE_THINKING: "gpt-oss:20b",
ROLE_FAST: "gemma4:latest",
ROLE_UTILITY: "llama3.2",
ROLE_EMBED: "nomic-embed-text",
}[role]
configured = getattr(settings, role_setting, None)
if configured:
return configured
# Embeddings must never silently fall back to a chat model (#62).
if role == ROLE_EMBED:
return role_default
legacy = getattr(settings, "OLLAMA_MODEL", None)
if legacy:
return legacy
return role_default
def ollama_embed_model() -> str:
return getattr(settings, "OLLAMA_EMBED_MODEL", ollama_model())
return ollama_model_for_role(ROLE_EMBED)
def ollama_llm_kwargs(**extra):
def ollama_num_ctx_for_role(role: str) -> int:
role = (role or ROLE_THINKING).lower()
if role == ROLE_FAST:
return int(getattr(settings, "OLLAMA_NUM_CTX_FAST", 8192) or 8192)
if role == ROLE_UTILITY:
return int(getattr(settings, "OLLAMA_NUM_CTX_UTILITY", 4096) or 4096)
return int(getattr(settings, "OLLAMA_NUM_CTX_THINKING", 16384) or 16384)
def resolve_chat_role(model_name: str | None) -> str:
"""Map FE mode selector (FAST / THINKING / …) to an Ollama role."""
if (model_name or "").upper() == "FAST":
return ROLE_FAST
return ROLE_THINKING
def ollama_llm_kwargs(role: str = ROLE_THINKING, **extra):
"""Keyword args for langchain_ollama.OllamaLLM / ChatOllama."""
kwargs = {"base_url": ollama_base_url(), "model": ollama_model()}
kwargs = {
"base_url": ollama_base_url(),
"model": ollama_model_for_role(role),
"num_ctx": ollama_num_ctx_for_role(role),
}
kwargs.update(extra)
return kwargs
+2
View File
@@ -223,7 +223,9 @@ class PromptSerializer(serializers.ModelSerializer):
"id",
"tokens_in",
"tokens_out",
"citations",
)
read_only_fields = ("citations",)
def _token_pair(self, obj):
cache = self.context.setdefault("_prompt_token_cache", {})
@@ -11,3 +11,17 @@ ASSISTANT_SYSTEM_PROMPT = (
"Your name evokes quiet, rest, silence, and stillness — "
"respond with calm clarity; keep answers focused and uncluttered."
)
GROUNDED_ANSWER_INSTRUCTIONS = (
"You have been given numbered live sources. Answer ONLY from those sources. "
"Cite source indexes inline like [1] or [2]. "
"If the sources do not settle the question, say so explicitly — do not fill "
"gaps from memory or training data. "
"When sources conflict, prefer the most recent dated source. "
"Never state a date, number, name, or event that does not appear in the sources."
)
RETRIEVAL_FAILED_MESSAGE = (
"I couldn't reach live sources to answer this accurately right now. "
"Please try again in a moment — I won't guess from outdated training data."
)
+8 -4
View File
@@ -1,20 +1,24 @@
from abc import ABC, abstractmethod
from abc import ABC
from langchain_ollama import OllamaLLM
from langchain_core.output_parsers import StrOutputParser
from chat_backend.ollama_config import ollama_llm_kwargs
from chat_backend.ollama_config import ROLE_UTILITY, ollama_llm_kwargs
class BaseService(ABC):
"""Abstract base class for LLM conversation services."""
def __init__(self, temperature=0.7):
def __init__(self, temperature=0.7, role: str = ROLE_UTILITY, **llm_extra):
self.role = role
self.llm = OllamaLLM(
**ollama_llm_kwargs(
role=role,
temperature=temperature,
top_k=50,
top_p=0.9,
repeat_penalty=1.1,
num_ctx=4096,
**llm_extra,
)
)
self.output_parser = StrOutputParser()
@@ -0,0 +1,123 @@
"""Apply always-on grounded retrieval for a chat turn (#62).
Shared by ``consumers`` and ``consumers_graph`` so both paths stay in sync.
"""
from __future__ import annotations
import logging
from dataclasses import dataclass, field
from typing import Any
from asgiref.sync import sync_to_async
from django.conf import settings
from chat_backend.services.assistant_identity import RETRIEVAL_FAILED_MESSAGE
from chat_backend.services.grounding_decider import (
GroundingDecision,
grounding_decider,
)
from chat_backend.services.llm_service import build_chat_service
from chat_backend.services.search import (
SearchUnavailable,
format_sources_block,
search_and_rank,
)
from chat_backend.services.search.base import SearchResult
logger = logging.getLogger(__name__)
@dataclass
class GroundedTurnResult:
"""Outcome of the grounding + optional retrieval step."""
generator: Any = None
error: dict | None = None
citations: list[dict] = field(default_factory=list)
grounded: bool = False
decision: GroundingDecision | None = None
model_name: str = ""
def _citations_from_results(results: list[SearchResult]) -> list[dict]:
return [r.to_citation(i) for i, r in enumerate(results, start=1)]
async def prepare_grounded_chat(
*,
message: str,
messages: list,
model_name: str | None,
conversation_id: int,
) -> GroundedTurnResult:
"""Decide grounding, retrieve, and return an AsyncLLMService generator.
When retrieval is required but every provider fails, returns an ``error``
dict instead of falling back to parametric generation (#62 AC).
"""
internet = getattr(settings, "ALLOW_INTERNET_ACCESS", False)
if not internet:
service = build_chat_service(model_name=model_name, grounded=False)
return GroundedTurnResult(
generator=service.generate_response(
messages, message, conversation_id
),
model_name=service.model_name,
)
decision = await grounding_decider.decide_async(message)
if not decision.needs_retrieval:
service = build_chat_service(model_name=model_name, grounded=False)
return GroundedTurnResult(
generator=service.generate_response(
messages, message, conversation_id
),
decision=decision,
model_name=service.model_name,
)
try:
results = await sync_to_async(search_and_rank, thread_sensitive=False)(
decision.queries or [message],
temporal=decision.temporal,
)
except SearchUnavailable as exc:
logger.warning(
"Grounded retrieval failed for %r (queries=%s): %s",
message,
decision.queries,
exc,
)
return GroundedTurnResult(
error={
"type": "error",
"code": "search_unavailable",
"content": RETRIEVAL_FAILED_MESSAGE,
},
decision=decision,
grounded=True,
)
sources_block = format_sources_block(results)
# Keep sources out of the mutable history list — AsyncLLMService injects
# them via {sources}. (Older code appended a HumanMessage; that double-
# rendered into the prompt.)
service = build_chat_service(
model_name=model_name,
grounded=True,
sources_block=sources_block,
)
return GroundedTurnResult(
generator=service.generate_response(messages, message, conversation_id),
citations=_citations_from_results(results),
grounded=True,
decision=decision,
model_name=service.model_name,
)
def citations_frame(citations: list[dict]) -> dict:
"""Versioned WS envelope agreed with FE progress/citations work."""
return {"v": 1, "type": "citations", "data": citations}
@@ -0,0 +1,203 @@
"""Grounding decision: retrieval-on-unless-unnecessary (#62).
Fails open — parse/timeout/exception ⇒ needs_retrieval=True. A deterministic
temporal-marker pre-pass forces retrieval regardless of the model.
"""
from __future__ import annotations
import json
import logging
import re
from dataclasses import dataclass, field
from langchain_core.prompts import ChatPromptTemplate
from chat_backend.ollama_config import ROLE_UTILITY
from chat_backend.services.base_service import BaseService
logger = logging.getLogger(__name__)
# Years at/after common small-model cutoffs force live retrieval.
_TRAINING_CUTOFF_YEAR = 2024
_TEMPORAL_PATTERNS = (
r"\blatest\b",
r"\bcurrent\b",
r"\btoday\b",
r"\bnow\b",
r"\bthis year\b",
r"\bthis week\b",
r"\bthis month\b",
r"\bbreaking\b",
r"\bright now\b",
r"\bas of\b",
r"\bdid\b.+\byet\b",
r"\bhave\b.+\byet\b",
r"\bwho won\b",
r"\bstock price\b",
r"\bweather\b",
rf"\b(?:19|20)\d{{2}}\b", # any year mention — keep broad; model still helps
)
_TEMPORAL_RE = re.compile("|".join(_TEMPORAL_PATTERNS), re.IGNORECASE)
_YEAR_RE = re.compile(r"\b((?:19|20)\d{2})\b")
@dataclass
class GroundingDecision:
needs_retrieval: bool
reason: str = ""
queries: list[str] = field(default_factory=list)
temporal: bool = False
source: str = "model" # prepass | model | fail_open
class GroundingDecider(BaseService):
def __init__(self):
super().__init__(temperature=0.0, role=ROLE_UTILITY)
self.prompt = ChatPromptTemplate.from_messages(
[
(
"system",
"""You decide whether a user message needs live web retrieval.
Return ONLY compact JSON with keys:
needs_retrieval (boolean),
reason (short string),
queries (array of 1-3 focused search queries).
Bias TOWARD retrieval. Set needs_retrieval=true for ANY question about:
- a real person, organisation, product, price, event, date, or statistic
- anything that can change over time or after a model training cutoff
- news, sports, celebrity, politics, weather, stock prices
Set needs_retrieval=false ONLY when the message is fully self-contained:
creative writing, math, code, chit-chat, or a pure follow-up on text already
in the conversation that needs no external facts.
When needs_retrieval=true, produce focused search queries (not the raw user
message). Example: "did Taylor Swift get married"
["Taylor Swift Travis Kelce wedding date", "Taylor Swift married 2026"].
""",
),
("human", "{prompt}"),
]
)
self.chain = self.prompt | self.llm
def temporal_prepass(self, prompt: str) -> GroundingDecision | None:
text = (prompt or "").strip()
if not text:
return GroundingDecision(
needs_retrieval=False,
reason="empty prompt",
source="prepass",
)
year_hits = [int(y) for y in _YEAR_RE.findall(text)]
forces = bool(_TEMPORAL_RE.search(text)) or any(
y >= _TRAINING_CUTOFF_YEAR for y in year_hits
)
if not forces:
return None
return GroundingDecision(
needs_retrieval=True,
reason="temporal marker / post-cutoff year",
queries=[text],
temporal=True,
source="prepass",
)
def _parse(self, raw: str, fallback_query: str) -> GroundingDecision:
text = (raw or "").strip()
# Strip markdown fences if the small model wraps JSON.
if text.startswith("```"):
text = re.sub(r"^```(?:json)?\s*", "", text)
text = re.sub(r"\s*```$", "", text)
try:
start = text.find("{")
end = text.rfind("}")
if start < 0 or end < 0:
raise ValueError("no JSON object")
payload = json.loads(text[start : end + 1])
except (ValueError, json.JSONDecodeError) as exc:
logger.warning("GroundingDecider parse fail (%s); failing open", exc)
return GroundingDecision(
needs_retrieval=True,
reason=f"unparseable decision: {exc}",
queries=[fallback_query],
temporal=True,
source="fail_open",
)
needs = bool(payload.get("needs_retrieval", True))
queries = payload.get("queries") or []
if not isinstance(queries, list):
queries = [str(queries)]
queries = [str(q).strip() for q in queries if str(q).strip()][:3]
if needs and not queries:
queries = [fallback_query]
return GroundingDecision(
needs_retrieval=needs,
reason=str(payload.get("reason") or ""),
queries=queries,
temporal=needs,
source="model",
)
async def decide_async(self, prompt: str) -> GroundingDecision:
pre = self.temporal_prepass(prompt)
# Even on temporal prepass, ask the model for better queries when possible.
try:
raw = await self.chain.ainvoke({"prompt": prompt})
if hasattr(raw, "content"):
raw = raw.content
decision = self._parse(str(raw), prompt)
except Exception as exc:
logger.warning("GroundingDecider LLM failed (%s); failing open", exc)
decision = GroundingDecision(
needs_retrieval=True,
reason=f"llm error: {exc}",
queries=[prompt],
temporal=True,
source="fail_open",
)
if pre and pre.needs_retrieval:
# Pre-pass wins on needs_retrieval; keep model queries when present.
return GroundingDecision(
needs_retrieval=True,
reason=pre.reason,
queries=decision.queries or pre.queries,
temporal=True,
source="prepass",
)
return decision
def decide(self, prompt: str) -> GroundingDecision:
pre = self.temporal_prepass(prompt)
try:
raw = self.chain.invoke({"prompt": prompt})
if hasattr(raw, "content"):
raw = raw.content
decision = self._parse(str(raw), prompt)
except Exception as exc:
logger.warning("GroundingDecider LLM failed (%s); failing open", exc)
decision = GroundingDecision(
needs_retrieval=True,
reason=f"llm error: {exc}",
queries=[prompt],
temporal=True,
source="fail_open",
)
if pre and pre.needs_retrieval:
return GroundingDecision(
needs_retrieval=True,
reason=pre.reason,
queries=decision.queries or pre.queries,
temporal=True,
source="prepass",
)
return decision
grounding_decider = GroundingDecider()
+128 -58
View File
@@ -1,28 +1,51 @@
from abc import ABC, abstractmethod
from typing import AsyncGenerator, Generator, Optional
# from langchain_community.llms import Ollama
from langchain_ollama import OllamaLLM
from langchain_core.output_parsers import StrOutputParser
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
from chat_backend.services.assistant_identity import ASSISTANT_SYSTEM_PROMPT
from chat_backend.ollama_config import (
ROLE_FAST,
ROLE_THINKING,
ollama_llm_kwargs,
ollama_model_for_role,
ollama_num_ctx_for_role,
resolve_chat_role,
)
from chat_backend.services.assistant_identity import (
ASSISTANT_SYSTEM_PROMPT,
GROUNDED_ANSWER_INSTRUCTIONS,
)
from chat_backend.services.prompt_budget import (
estimate_tokens,
format_history,
window_history,
)
class LLMService(ABC):
"""Abstract base class for LLM conversation services."""
def __init__(self):
def __init__(
self,
role: str = ROLE_THINKING,
temperature: float = 0.7,
grounded: bool = False,
):
self.role = role
self.grounded = grounded
self.model_name = ollama_model_for_role(role)
self.num_ctx = ollama_num_ctx_for_role(role)
self.llm = OllamaLLM(
**ollama_llm_kwargs(
temperature=0.7,
role=role,
temperature=temperature,
top_k=50,
top_p=0.9,
repeat_penalty=1.1,
num_ctx=4096,
)
)
self.output_parser = StrOutputParser()
@@ -45,8 +68,8 @@ class LLMService(ABC):
class SyncLLMService(LLMService):
"""Synchronous LLM conversation service."""
def __init__(self):
super().__init__()
def __init__(self, role: str = ROLE_THINKING, temperature: float = 0.7):
super().__init__(role=role, temperature=temperature)
self._setup_chain()
def _setup_chain(self):
@@ -85,36 +108,49 @@ class SyncLLMService(LLMService):
class AsyncLLMService(LLMService):
"""Asynchronous LLM conversation service."""
def __init__(self):
super().__init__()
def __init__(
self,
role: str = ROLE_THINKING,
temperature: float = 0.7,
grounded: bool = False,
sources_block: str = "",
):
super().__init__(role=role, temperature=temperature, grounded=grounded)
self.sources_block = sources_block or ""
self._setup_chain()
def _setup_chain(self):
"""Setup the conversation chain."""
template = f"""{ASSISTANT_SYSTEM_PROMPT}
"""Single history window + optional grounded sources (#62 Phase 3)."""
grounded_block = ""
if self.grounded:
grounded_block = (
f"\n\n{GROUNDED_ANSWER_INSTRUCTIONS}\n\n"
f"Live sources:\n{{sources}}\n"
)
Continue this conversation while maintaining context by providing a single helpful response.
Current context: {{context}}
Last 3 messages:
{{recent_history}}
Latest message: {{query}}
Instructions:
- Carefully maintain all established context
- If referencing previous elements (like stories), preserve all details
- When asked to modify something, identify what's being modified
Response:"""
template = f"""{ASSISTANT_SYSTEM_PROMPT}
{grounded_block}
Continue this conversation while maintaining context by providing a single helpful response.
Conversation history:
{{history}}
Latest message: {{query}}
Instructions:
- Carefully maintain all established context
- If referencing previous elements (like stories), preserve all details
- When asked to modify something, identify what's being modified
Response:"""
self.prompt = ChatPromptTemplate.from_template(template)
self.conversation_chain = (
{
"context":lambda x: x["conversation"],
"recent_history":lambda x: x['recent_conversation'],
"history": lambda x: x["history"],
"query": lambda x: x["query"],
"sources": lambda x: x.get("sources", ""),
}
| self.prompt
| self.llm
@@ -122,39 +158,73 @@ class AsyncLLMService(LLMService):
# final GenerationChunk.generation_info; the parser would drop it.
)
async def _format_history(self, conversation: list) -> str:
"""Async version of format conversation history."""
# prompts = list(
# await Prompt.objects.filter(conversation_id=conversation_id)
# .order_by("created")
# )
# 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])
async def _get_recent_messages(self, conversation: list) -> str:
"""Async version of format conversation history."""
# prompts = list(
# await Prompt.objects.filter(conversation_id=conversation_id)
# .order_by("created")
# [-6:]
# )
# 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])
async def generate_response(
self, conversation: Conversation, query: str, conversation_id: int, **kwargs
self,
conversation,
query: str,
conversation_id: int,
**kwargs,
) -> AsyncGenerator[str, None]:
"""Generate response with async streaming support."""
"""Generate response with async streaming support.
``conversation`` is the LangChain message list for this turn (not a
Django Conversation row). History is serialised exactly once and
trimmed oldest-first under the role's ``num_ctx`` budget. The sources
block (when grounded) is never trimmed.
"""
sources = self.sources_block or kwargs.get("sources_block", "") or ""
reserved = (
estimate_tokens(ASSISTANT_SYSTEM_PROMPT)
+ estimate_tokens(GROUNDED_ANSWER_INSTRUCTIONS if self.grounded else "")
+ estimate_tokens(sources)
+ estimate_tokens(query)
+ 256 # response headroom / instructions
)
# Leave ~40% of ctx for the completion.
history_budget = max(512, int(self.num_ctx * 0.55) - reserved)
# Drop any prior "Search Results:" blobs from history — sources are
# passed separately now so we don't double-inject.
clean = [
m
for m in conversation
if not (
getattr(m, "type", "") == "human"
and str(getattr(m, "content", "")).startswith("Search Results:")
)
and not (
getattr(m, "type", "") == "human"
and str(getattr(m, "content", "")).startswith("Live sources:")
)
]
# Exclude the latest user turn from history (it's in {query}).
prior = clean[:-1] if clean else []
windowed = window_history(
prior, budget_tokens=history_budget, reserved_tokens=0
)
history_text = format_history(windowed)
chain_input = {
"query": query,
"conversation": await self._format_history(conversation),
"recent_conversation": await self._get_recent_messages(conversation[-6:])}
"query": query,
"history": history_text,
"sources": sources,
}
async for chunk in self.conversation_chain.astream(chain_input):
yield chunk
def build_chat_service(
*,
model_name: str | None,
grounded: bool = False,
sources_block: str = "",
) -> AsyncLLMService:
"""Factory: FE mode → role, factual turns → low temperature."""
role = resolve_chat_role(model_name)
temperature = 0.3 if grounded else 0.7
return AsyncLLMService(
role=role,
temperature=temperature,
grounded=grounded,
sources_block=sources_block,
)
@@ -3,7 +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
from chat_backend.ollama_config import ROLE_UTILITY, ollama_llm_kwargs
class ModerationLabel(Enum):
@@ -18,9 +18,10 @@ class ModerationClassifier(BaseService):
"""
def __init__(self):
super().__init__(temperature=0.1)
super().__init__(temperature=0.1, role=ROLE_UTILITY)
self.llm = OllamaLLM(
**ollama_llm_kwargs(
role=ROLE_UTILITY,
temperature=0.1, # Very low for strict moderation
top_k=10,
num_ctx=2048,
@@ -0,0 +1,63 @@
"""Prompt budgeting helpers (#62 Phase 3).
Approximate token counts with a chars/4 heuristic. Trim oldest history first;
never truncate the system prompt or the retrieved-sources block.
"""
from __future__ import annotations
from typing import Sequence
def estimate_tokens(text: str) -> int:
if not text:
return 0
return max(1, (len(text) + 3) // 4)
def window_history(
messages: Sequence,
*,
budget_tokens: int,
reserved_tokens: int = 0,
max_messages: int = 24,
) -> list:
"""Keep the newest messages that fit under ``budget_tokens - reserved``.
``messages`` are LangChain BaseMessage-like objects (``.type``, ``.text`` /
``.content``).
"""
if not messages:
return []
usable = max(0, budget_tokens - reserved_tokens)
selected: list = []
used = 0
for message in reversed(list(messages)[-max_messages:]):
text = getattr(message, "text", None)
if callable(text):
# property that looks callable in some versions — read content
text = getattr(message, "content", "")
if text is None:
text = getattr(message, "content", "") or ""
cost = estimate_tokens(str(text)) + 4 # role overhead
if selected and used + cost > usable:
break
selected.append(message)
used += cost
selected.reverse()
return selected
def format_message_line(message) -> str:
role = "User" if getattr(message, "type", "") == "human" else "AI"
text = getattr(message, "text", None)
if text is None or (callable(text) and not isinstance(text, str)):
text = getattr(message, "content", "") or ""
# BaseMessage.text is a property returning str; prefer it when string.
if not isinstance(text, str):
text = getattr(message, "content", "") or ""
return f"{role}: {text}"
def format_history(messages: Sequence) -> str:
return "\n".join(format_message_line(m) for m in messages)
+37 -1
View File
@@ -25,7 +25,14 @@ from chat_backend.models import Conversation, Prompt, DocumentWorkspace, Documen
from pathlib import Path
from chat_backend.services.base_service import BaseService
from chat_backend.services.assistant_identity import ASSISTANT_SYSTEM_PROMPT
from chat_backend.ollama_config import ollama_embeddings_kwargs
from chat_backend.ollama_config import ollama_embed_model, ollama_embeddings_kwargs
import logging
logger = logging.getLogger(__name__)
class EmbeddingDimensionMismatch(RuntimeError):
"""Persisted Chroma collection dim != configured embedding model (#62)."""
@database_sync_to_async
@@ -54,6 +61,7 @@ class RAGService(BaseService):
chunk_size=1000, chunk_overlap=200
)
self.vector_store = self._initialize_vector_store()
self._assert_embedding_dimensions()
# Supported file types and their loaders
self.loader_mapping = {
@@ -75,6 +83,34 @@ class RAGService(BaseService):
)
return vector_store
def _assert_embedding_dimensions(self) -> None:
"""Refuse mismatched collections loudly (#62)."""
try:
collection = self.vector_store._collection
count = collection.count()
except Exception as exc:
logger.warning("Could not inspect Chroma collection: %s", exc)
return
if not count:
return
try:
peek = collection.peek(limit=1)
embeddings = peek.get("embeddings") if isinstance(peek, dict) else None
if not embeddings:
return
stored_dim = len(embeddings[0])
probe = self.embedding_model.embed_query("dimension-check")
expected_dim = len(probe)
except Exception as exc:
logger.warning("Embedding dimension probe failed: %s", exc)
return
if stored_dim != expected_dim:
raise EmbeddingDimensionMismatch(
f"Chroma collection embedding dim is {stored_dim} but "
f"OLLAMA_EMBED_MODEL={ollama_embed_model()!r} produces "
f"{expected_dim}. Run: python manage.py reindex_embeddings"
)
def clear_vector_store(self):
"""Clear all vectors from the store"""
self.vector_store.delete_collection()
@@ -0,0 +1,22 @@
"""Structured web-search providers (#62).
Providers return typed :class:`SearchResult` rows — never a flat concatenated
string. The facade in :mod:`chat_backend.services.search.service` picks the
configured primary provider and fails over to the secondary.
"""
from chat_backend.services.search.base import SearchResult
from chat_backend.services.search.service import (
SearchUnavailable,
format_sources_block,
get_search_service,
search_and_rank,
)
__all__ = [
"SearchResult",
"SearchUnavailable",
"format_sources_block",
"get_search_service",
"search_and_rank",
]
@@ -0,0 +1,36 @@
"""Search provider protocol and result dataclass."""
from __future__ import annotations
from dataclasses import asdict, dataclass
from typing import Optional, Protocol, runtime_checkable
@dataclass(frozen=True)
class SearchResult:
title: str
url: str
snippet: str
published_at: Optional[str] = None
rank: int = 0
provider: str = ""
def to_citation(self, index: int) -> dict:
return {
"index": index,
"title": self.title,
"url": self.url,
"published_at": self.published_at,
}
def to_dict(self) -> dict:
return asdict(self)
@runtime_checkable
class SearchProvider(Protocol):
name: str
def search(self, query: str, *, max_results: int = 8) -> list[SearchResult]:
"""Return structured results for ``query``. Raise on hard failure."""
...
@@ -0,0 +1,46 @@
"""DuckDuckGo (ddgs) search provider — failover for SearxNG (#62)."""
from __future__ import annotations
import logging
from typing import Any
from chat_backend.services.search.base import SearchResult
logger = logging.getLogger(__name__)
class DDGSProvider:
name = "ddgs"
def search(self, query: str, *, max_results: int = 8) -> list[SearchResult]:
try:
from ddgs import DDGS
except ImportError as exc: # pragma: no cover - dependency is declared
raise RuntimeError("ddgs package is not installed") from exc
try:
with DDGS() as ddgs:
raw: list[dict[str, Any]] = list(
ddgs.text(query, max_results=max_results)
)
except Exception as exc:
logger.warning("DDGS search failed for %r: %s", query, exc)
raise RuntimeError(f"DDGS unreachable: {exc}") from exc
results: list[SearchResult] = []
for idx, item in enumerate(raw or []):
url = (item.get("href") or item.get("link") or item.get("url") or "").strip()
if not url:
continue
results.append(
SearchResult(
title=(item.get("title") or url).strip(),
url=url,
snippet=(item.get("body") or item.get("snippet") or "").strip(),
published_at=item.get("date") or item.get("published") or None,
rank=idx,
provider=self.name,
)
)
return results
@@ -0,0 +1,128 @@
"""Rank, dedupe, and rumour-filter search results (#62)."""
from __future__ import annotations
import re
from datetime import datetime
from typing import Iterable
from urllib.parse import urlparse
from chat_backend.services.search.base import SearchResult
_RUMOUR_MARKERS = (
"rumor",
"rumour",
"speculation",
"ai-generated",
"ai generated",
"blind item",
"fake",
"allegedly",
"unconfirmed",
)
_DATE_FORMATS = (
"%Y-%m-%d",
"%Y-%m-%dT%H:%M:%S",
"%Y-%m-%dT%H:%M:%SZ",
"%Y-%m-%dT%H:%M:%S%z",
"%b %d, %Y",
"%B %d, %Y",
"%d %b %Y",
"%d %B %Y",
)
def registrable_domain(url: str) -> str:
host = (urlparse(url).hostname or "").lower()
if host.startswith("www."):
host = host[4:]
parts = host.split(".")
if len(parts) >= 2:
return ".".join(parts[-2:])
return host
def parse_published_at(value: str | None) -> datetime | None:
if not value:
return None
text = value.strip()
if not text:
return None
# Prefer ISO-ish prefixes.
for fmt in _DATE_FORMATS:
try:
return datetime.strptime(text[: len(fmt) + 8], fmt)
except ValueError:
continue
iso = text.replace("Z", "+00:00")
try:
return datetime.fromisoformat(iso)
except ValueError:
return None
def is_rumour_heavy(result: SearchResult) -> bool:
haystack = f"{result.title} {result.snippet}".lower()
hits = sum(1 for marker in _RUMOUR_MARKERS if marker in haystack)
return hits >= 2 or ("ai-generated" in haystack and "fake" in haystack)
def rank_and_dedupe(
results: Iterable[SearchResult],
*,
temporal: bool = False,
max_results: int = 6,
) -> list[SearchResult]:
"""Deduplicate by domain, drop rumour-heavy rows when alternatives exist."""
seen_domains: set[str] = set()
kept: list[SearchResult] = []
rumour_bucket: list[SearchResult] = []
for result in results:
domain = registrable_domain(result.url)
if not domain or domain in seen_domains:
continue
seen_domains.add(domain)
if is_rumour_heavy(result):
rumour_bucket.append(result)
else:
kept.append(result)
# Only use rumour-heavy rows if we have nothing better.
if not kept and rumour_bucket:
kept = rumour_bucket
def sort_key(item: SearchResult):
published = parse_published_at(item.published_at)
# Prefer dated + recent when temporal; otherwise keep provider rank.
if temporal:
# Newer first; undated last.
stamp = published.timestamp() if published else float("-inf")
return (-stamp, item.rank)
has_date = 0 if published else 1
return (has_date, item.rank)
kept.sort(key=sort_key)
return [
SearchResult(
title=r.title,
url=r.url,
snippet=r.snippet,
published_at=r.published_at,
rank=i,
provider=r.provider,
)
for i, r in enumerate(kept[:max_results])
]
_HOST_RE = re.compile(r"^https?://([^/]+)", re.I)
def display_host(url: str) -> str:
match = _HOST_RE.match(url or "")
if not match:
return url or ""
host = match.group(1).lower()
return host[4:] if host.startswith("www.") else host
@@ -0,0 +1,83 @@
"""SearxNG search provider (#62).
Hits a self-hosted SearxNG instance's JSON API. Deterministic, no third-party
rate limits, and returns per-result title/url/snippet/publishedDate.
"""
from __future__ import annotations
import logging
from typing import Any
from urllib.parse import urljoin
import requests
from django.conf import settings
from chat_backend.services.search.base import SearchResult
logger = logging.getLogger(__name__)
class SearxNGProvider:
name = "searxng"
def __init__(
self,
base_url: str | None = None,
timeout: float | None = None,
):
self.base_url = (
base_url
or getattr(settings, "SEARXNG_BASE_URL", "http://127.0.0.1:8080")
).rstrip("/")
self.timeout = float(
timeout
if timeout is not None
else getattr(settings, "SEARXNG_TIMEOUT_SECONDS", 8)
)
def search(self, query: str, *, max_results: int = 8) -> list[SearchResult]:
endpoint = urljoin(self.base_url + "/", "search")
params = {
"q": query,
"format": "json",
"language": "en",
}
try:
response = requests.get(
endpoint,
params=params,
timeout=self.timeout,
headers={"Accept": "application/json"},
)
response.raise_for_status()
payload: dict[str, Any] = response.json()
except requests.RequestException as exc:
logger.warning("SearxNG search failed for %r: %s", query, exc)
raise RuntimeError(f"SearxNG unreachable: {exc}") from exc
except ValueError as exc:
raise RuntimeError(f"SearxNG returned non-JSON: {exc}") from exc
results: list[SearchResult] = []
for idx, item in enumerate(payload.get("results") or []):
url = (item.get("url") or "").strip()
if not url:
continue
results.append(
SearchResult(
title=(item.get("title") or url).strip(),
url=url,
snippet=(item.get("content") or item.get("snippet") or "").strip(),
published_at=(
item.get("publishedDate")
or item.get("published_at")
or item.get("pubdate")
or None
),
rank=idx,
provider=self.name,
)
)
if len(results) >= max_results:
break
return results
@@ -0,0 +1,145 @@
"""Search facade: primary + failover providers, concurrent multi-query (#62)."""
from __future__ import annotations
import logging
from concurrent.futures import ThreadPoolExecutor, as_completed
from typing import Iterable
from django.conf import settings
from chat_backend.services.search.base import SearchProvider, SearchResult
from chat_backend.services.search.ddgs_provider import DDGSProvider
from chat_backend.services.search.ranking import display_host, rank_and_dedupe
from chat_backend.services.search.searxng import SearxNGProvider
logger = logging.getLogger(__name__)
class SearchUnavailable(RuntimeError):
"""Raised when every configured provider fails for a query set."""
def _build_provider(name: str) -> SearchProvider:
key = (name or "").lower()
if key in {"searxng", "searx"}:
return SearxNGProvider()
if key in {"ddgs", "duckduckgo", "ddg"}:
return DDGSProvider()
raise ValueError(f"Unknown search provider: {name!r}")
class SearchService:
def __init__(
self,
primary: SearchProvider | None = None,
failover: SearchProvider | None = None,
):
primary_name = getattr(settings, "SEARCH_PROVIDER", "searxng")
failover_name = getattr(settings, "SEARCH_FAILOVER_PROVIDER", "ddgs")
self.primary = primary or _build_provider(primary_name)
# Avoid wiring the same provider twice.
if failover is not None:
self.failover = failover
elif failover_name and failover_name.lower() != getattr(
self.primary, "name", ""
):
try:
self.failover = _build_provider(failover_name)
except ValueError:
self.failover = None
else:
self.failover = None
def search_one(self, query: str, *, max_results: int = 8) -> list[SearchResult]:
errors: list[str] = []
for provider in (self.primary, self.failover):
if provider is None:
continue
try:
return provider.search(query, max_results=max_results)
except Exception as exc:
errors.append(f"{getattr(provider, 'name', provider)}: {exc}")
logger.warning(
"Search provider %s failed for %r: %s",
getattr(provider, "name", provider),
query,
exc,
)
raise SearchUnavailable(
f"All search providers failed for {query!r}: {'; '.join(errors)}"
)
def search_many(
self,
queries: Iterable[str],
*,
max_results_per_query: int = 6,
temporal: bool = False,
max_results: int = 6,
) -> list[SearchResult]:
cleaned = [q.strip() for q in queries if q and str(q).strip()]
if not cleaned:
raise SearchUnavailable("No search queries provided")
collected: list[SearchResult] = []
failures = 0
with ThreadPoolExecutor(max_workers=min(3, len(cleaned))) as pool:
futures = {
pool.submit(
self.search_one, query, max_results=max_results_per_query
): query
for query in cleaned[:3]
}
for future in as_completed(futures):
query = futures[future]
try:
collected.extend(future.result())
except SearchUnavailable as exc:
failures += 1
logger.warning("Query %r exhausted providers: %s", query, exc)
if not collected:
raise SearchUnavailable(
f"All {failures} search queries failed; no results available"
)
return rank_and_dedupe(
collected, temporal=temporal, max_results=max_results
)
_service: SearchService | None = None
def get_search_service() -> SearchService:
global _service
if _service is None:
_service = SearchService()
return _service
def search_and_rank(
queries: Iterable[str],
*,
temporal: bool = False,
max_results: int = 6,
) -> list[SearchResult]:
return get_search_service().search_many(
queries, temporal=temporal, max_results=max_results
)
def format_sources_block(results: list[SearchResult]) -> str:
"""Numbered, dated, delimited context block for the LLM prompt."""
if not results:
return ""
lines: list[str] = []
for i, result in enumerate(results, start=1):
host = display_host(result.url)
date = result.published_at or "undated"
lines.append(f'[{i}] "{result.title}"{host}{date}')
if result.snippet:
lines.append(f" {result.snippet}")
lines.append(f" URL: {result.url}")
lines.append("")
return "\n".join(lines).rstrip()
@@ -3,7 +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
from chat_backend.ollama_config import ROLE_UTILITY, ollama_llm_kwargs
class TitleGenerator:
@@ -14,6 +14,7 @@ class TitleGenerator:
def __init__(self):
self.llm = OllamaLLM(
**ollama_llm_kwargs(
role=ROLE_UTILITY,
temperature=0.5, # Slightly creative but not too random
top_k=20,
num_ctx=2048, # Shorter context needed for titles
+72 -27
View File
@@ -422,58 +422,103 @@ class GraphNodeTestCase(TransactionTestCase):
self.assertEqual(result["response_generator"], "generator")
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"
with mock.patch(
"chat_backend.consumers_graph.prepare_grounded_chat"
) as prepare:
from chat_backend.services.grounded_chat import GroundedTurnResult
prepare.return_value = GroundedTurnResult(
generator="generator", model_name="gpt-oss:20b"
)
result = await consumers_graph.generation_node(self._state())
self.assertEqual(result["response_generator"], "generator")
service.return_value.generate_response.assert_called_once()
prepare.assert_called_once()
@override_settings(ALLOW_INTERNET_ACCESS=True)
async def test_search_prompts_append_web_results(self):
async def test_search_prompts_use_grounded_chat(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)
with mock.patch(
"chat_backend.consumers_graph.prepare_grounded_chat"
) as prepare:
from chat_backend.services.grounded_chat import GroundedTurnResult
self.assertIn("Search Results: top result", state["messages"][-1].content)
prepare.return_value = GroundedTurnResult(
generator="generator",
citations=[
{
"index": 1,
"title": "T",
"url": "https://example.com",
"published_at": None,
}
],
grounded=True,
model_name="gpt-oss:20b",
)
result = await consumers_graph.generation_node(state)
prepare.assert_called_once()
self.assertEqual(result["response_generator"], "generator")
self.assertEqual(result["citations"][0]["url"], "https://example.com")
@override_settings(ALLOW_INTERNET_ACCESS=True)
async def test_fast_model_skips_web_search(self):
async def test_fast_model_still_runs_grounding(self):
"""FAST must not skip retrieval (#62) — it only picks a smaller model."""
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)
with mock.patch(
"chat_backend.consumers_graph.prepare_grounded_chat"
) as prepare:
from chat_backend.services.grounded_chat import GroundedTurnResult
search.assert_not_called()
self.assertEqual(len(state["messages"]), 1)
prepare.return_value = GroundedTurnResult(
generator="generator", model_name="gemma4:latest"
)
await consumers_graph.generation_node(state)
prepare.assert_called_once()
self.assertEqual(prepare.call_args.kwargs["model_name"], "FAST")
@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)
with mock.patch(
"chat_backend.consumers_graph.prepare_grounded_chat"
) as prepare:
from chat_backend.services.grounded_chat import GroundedTurnResult
search.assert_not_called()
prepare.return_value = GroundedTurnResult(
generator="generator", model_name="gpt-oss:20b"
)
await consumers_graph.generation_node(state)
# prepare_grounded_chat still runs; inside it skips providers when
# ALLOW_INTERNET_ACCESS is False.
prepare.assert_called_once()
@override_settings(ALLOW_INTERNET_ACCESS=True)
async def test_search_failures_fall_back_to_plain_chat(self):
async def test_search_failures_surface_error_not_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)
with mock.patch(
"chat_backend.consumers_graph.prepare_grounded_chat"
) as prepare:
from chat_backend.services.grounded_chat import GroundedTurnResult
self.assertEqual(result["response_generator"], "generator")
self.assertEqual(len(state["messages"]), 1)
prepare.return_value = GroundedTurnResult(
error={
"type": "error",
"code": "search_unavailable",
"content": "couldn't reach live sources",
},
grounded=True,
)
result = await consumers_graph.generation_node(state)
self.assertEqual(result["response_generator"]["code"], "search_unavailable")
class WebSocketRoutingTestCase(TransactionTestCase):
@@ -0,0 +1,180 @@
"""Unit tests for grounding + search layer (#62 Phases 13)."""
from __future__ import annotations
from django.test import SimpleTestCase, override_settings
from unittest import mock
from chat_backend.services.grounding_decider import GroundingDecider, GroundingDecision
from chat_backend.services.search.base import SearchResult
from chat_backend.services.search.ranking import (
is_rumour_heavy,
rank_and_dedupe,
registrable_domain,
)
from chat_backend.services.search.service import (
SearchService,
SearchUnavailable,
format_sources_block,
)
from chat_backend.ollama_config import (
ROLE_EMBED,
ROLE_FAST,
ROLE_THINKING,
ROLE_UTILITY,
ollama_model_for_role,
resolve_chat_role,
)
class OllamaConfigRoleTestCase(SimpleTestCase):
@override_settings(
OLLAMA_MODEL="legacy-model",
OLLAMA_MODEL_THINKING="think-model",
OLLAMA_MODEL_FAST="fast-model",
OLLAMA_MODEL_UTILITY="util-model",
OLLAMA_EMBED_MODEL="nomic-embed-text",
)
def test_role_resolution(self):
self.assertEqual(ollama_model_for_role(ROLE_THINKING), "think-model")
self.assertEqual(ollama_model_for_role(ROLE_FAST), "fast-model")
self.assertEqual(ollama_model_for_role(ROLE_UTILITY), "util-model")
self.assertEqual(ollama_model_for_role(ROLE_EMBED), "nomic-embed-text")
@override_settings(
OLLAMA_MODEL="legacy-model",
OLLAMA_MODEL_THINKING="",
OLLAMA_MODEL_FAST="",
OLLAMA_MODEL_UTILITY="",
OLLAMA_EMBED_MODEL="",
)
def test_embed_never_falls_back_to_chat_model(self):
# Empty embed setting → hard default, not OLLAMA_MODEL.
self.assertEqual(ollama_model_for_role(ROLE_EMBED), "nomic-embed-text")
def test_resolve_chat_role(self):
self.assertEqual(resolve_chat_role("FAST"), ROLE_FAST)
self.assertEqual(resolve_chat_role("THINKING"), ROLE_THINKING)
self.assertEqual(resolve_chat_role(None), ROLE_THINKING)
class GroundingPrepassTestCase(SimpleTestCase):
def setUp(self):
self.decider = GroundingDecider.__new__(GroundingDecider)
def test_temporal_markers_force_retrieval(self):
cases = [
"did Taylor Swift get married yet",
"What is the latest news on AI?",
"Who won the Super Bowl this year?",
"current stock price of Apple",
"What happened in 2025?",
]
for prompt in cases:
with self.subTest(prompt=prompt):
decision = self.decider.temporal_prepass(prompt)
self.assertIsNotNone(decision)
self.assertTrue(decision.needs_retrieval)
def test_creative_prompt_does_not_force(self):
decision = self.decider.temporal_prepass("Write a poem about cats")
self.assertIsNone(decision)
def test_parse_failure_fails_open(self):
decision = self.decider._parse("NOT JSON", "fallback query")
self.assertTrue(decision.needs_retrieval)
self.assertEqual(decision.source, "fail_open")
self.assertEqual(decision.queries, ["fallback query"])
def test_decide_fail_open_on_llm_error(self):
self.decider.chain = mock.Mock()
self.decider.chain.invoke.side_effect = RuntimeError("ollama down")
decision = self.decider.decide("Is the sky blue?")
self.assertTrue(decision.needs_retrieval)
self.assertEqual(decision.source, "fail_open")
class SearchRankingTestCase(SimpleTestCase):
def test_dedupe_by_domain(self):
results = [
SearchResult("A", "https://www.people.com/a", "married", "2026-07-03", 0),
SearchResult("B", "https://people.com/b", "also", "2026-07-02", 1),
SearchResult("C", "https://bbc.com/c", "ok", "2026-07-01", 2),
]
ranked = rank_and_dedupe(results, temporal=True)
domains = {registrable_domain(r.url) for r in ranked}
self.assertEqual(domains, {"people.com", "bbc.com"})
def test_rumour_heavy_dropped_when_alternatives_exist(self):
clean = SearchResult(
"Married",
"https://people.com/wedding",
"Taylor Swift and Travis Kelce married July 3",
"2026-07-03",
0,
)
poison = SearchResult(
"Rumors",
"https://gossip.com/fake",
"fake AI-generated photos and speculation and blind item",
None,
1,
)
self.assertTrue(is_rumour_heavy(poison))
ranked = rank_and_dedupe([poison, clean], temporal=True)
self.assertEqual(len(ranked), 1)
self.assertEqual(ranked[0].url, clean.url)
def test_format_sources_block_is_numbered(self):
block = format_sources_block(
[
SearchResult(
"Title",
"https://example.com/x",
"Snippet here",
"2026-07-03",
0,
)
]
)
self.assertIn('[1] "Title"', block)
self.assertIn("2026-07-03", block)
self.assertIn("URL: https://example.com/x", block)
self.assertNotIn("Search Results:", block)
class FakeProvider:
def __init__(self, name, results=None, error=None):
self.name = name
self.results = results or []
self.error = error
self.calls = 0
def search(self, query, *, max_results=8):
self.calls += 1
if self.error:
raise self.error
return list(self.results)
class SearchFailoverTestCase(SimpleTestCase):
def test_failover_when_primary_raises(self):
primary = FakeProvider("searxng", error=RuntimeError("down"))
failover = FakeProvider(
"ddgs",
results=[
SearchResult("T", "https://a.com", "s", None, 0, "ddgs"),
],
)
service = SearchService(primary=primary, failover=failover)
results = service.search_one("q")
self.assertEqual(primary.calls, 1)
self.assertEqual(failover.calls, 1)
self.assertEqual(results[0].provider, "ddgs")
def test_all_providers_fail_raises(self):
primary = FakeProvider("searxng", error=RuntimeError("down"))
failover = FakeProvider("ddgs", error=RuntimeError("also down"))
service = SearchService(primary=primary, failover=failover)
with self.assertRaises(SearchUnavailable):
service.search_one("q")
+32 -18
View File
@@ -2,6 +2,7 @@ from django.test import SimpleTestCase
from langchain_core.messages import AIMessage, HumanMessage
from chat_backend.services.llm_service import AsyncLLMService, SyncLLMService
from chat_backend.services.prompt_budget import format_history, window_history
from .fakes import FakeChain
@@ -18,13 +19,6 @@ 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!"])
@@ -37,7 +31,7 @@ class AsyncLLMServiceTestCase(SimpleTestCase):
self.assertEqual("".join(chunks), "Hello!")
async def test_generate_response_sends_full_and_recent_history(self):
async def test_generate_response_sends_single_history_window(self):
self.service.conversation_chain = FakeChain(chunks=["ok"])
messages = conversation(4) # 8 messages
@@ -46,19 +40,39 @@ class AsyncLLMServiceTestCase(SimpleTestCase):
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"))
# Latest user turn is in {query}; history is prior turns only, once.
self.assertIn("history", payload)
self.assertNotIn("recent_conversation", payload)
self.assertNotIn("conversation", payload)
# 8 messages → drop last (query) → 7 prior lines max in window.
self.assertLessEqual(len(payload["history"].splitlines()), 7)
async def test_grounded_service_includes_sources(self):
service = AsyncLLMService(grounded=True, sources_block='[1] "T" — x.com — undated')
service.conversation_chain = FakeChain(chunks=["ok"])
async for _ in service.generate_response([], "q", 1):
pass
payload = service.conversation_chain.calls[0]
self.assertIn("[1]", payload["sources"])
class SyncLLMServiceTestCase(SimpleTestCase):
def test_generate_response_streams_chunks(self):
service = SyncLLMService()
service.conversation_chain = FakeChain(chunks=["one ", "two"])
def test_constructs(self):
self.assertIsNotNone(SyncLLMService())
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}]
class PromptBudgetTestCase(SimpleTestCase):
def test_window_history_drops_oldest_first(self):
messages = [HumanMessage(content="x" * 40) for _ in range(10)]
kept = window_history(messages, budget_tokens=30, reserved_tokens=0)
self.assertLess(len(kept), 10)
# Newest messages survive.
self.assertEqual(kept[-1].content, messages[-1].content)
def test_format_history_labels_speakers(self):
text = format_history(
[HumanMessage(content="hello"), AIMessage(content="hi")]
)
self.assertEqual(text, "User: hello\nAI: hi")
+11 -4
View File
@@ -137,7 +137,11 @@ class LastDayOfMonthTestCase(SimpleTestCase):
@override_settings(
OLLAMA_BASE_URL="http://10.0.0.128:11434",
OLLAMA_MODEL="llama3.2",
OLLAMA_MODEL_THINKING="llama3.2",
OLLAMA_MODEL_FAST="gemma4:latest",
OLLAMA_MODEL_UTILITY="llama3.2",
OLLAMA_EMBED_MODEL="nomic-embed-text",
OLLAMA_NUM_CTX_THINKING=16384,
)
class OllamaConfigTestCase(SimpleTestCase):
def test_reads_settings(self):
@@ -179,15 +183,18 @@ class OllamaConfigFallbackTestCase(SimpleTestCase):
del settings.OLLAMA_BASE_URL
del settings.OLLAMA_MODEL
del settings.OLLAMA_EMBED_MODEL
if hasattr(settings, "OLLAMA_MODEL_THINKING"):
del settings.OLLAMA_MODEL_THINKING
self.assertEqual(ollama_base_url(), "http://127.0.0.1:11434")
self.assertEqual(ollama_model(), "llama3.2")
self.assertEqual(ollama_embed_model(), "llama3.2")
self.assertEqual(ollama_model(), "gpt-oss:20b")
# Embeddings never fall back to a chat model (#62).
self.assertEqual(ollama_embed_model(), "nomic-embed-text")
def test_embed_model_falls_back_to_chat_model(self):
def test_embed_model_does_not_fall_back_to_chat_model(self):
with override_settings(OLLAMA_MODEL="llama3.2"):
del settings.OLLAMA_EMBED_MODEL
self.assertEqual(ollama_embed_model(), "llama3.2")
self.assertEqual(ollama_embed_model(), "nomic-embed-text")
class UserPromptGuardTestCase(SimpleTestCase):
+23 -5
View File
@@ -148,14 +148,23 @@ CORS_ALLOWED_ORIGINS = with_capacitor_webview_origins(
# 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
# Role-scoped models (#62). Dev and prod share the same defaults; override per
# role. ``OLLAMA_MODEL`` remains a fallback for THINKING so existing secrets
# keep working until role vars are set. Embeddings never fall back to a chat
# model.
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",
) or ("llama3.2" if not DEBUG else "gpt-oss:20b")
OLLAMA_EMBED_MODEL = env("OLLAMA_EMBED_MODEL", OLLAMA_MODEL) or OLLAMA_MODEL
OLLAMA_MODEL = env("OLLAMA_MODEL", "gpt-oss:20b") or "gpt-oss:20b"
OLLAMA_MODEL_THINKING = env("OLLAMA_MODEL_THINKING", "") or OLLAMA_MODEL
OLLAMA_MODEL_FAST = env("OLLAMA_MODEL_FAST", "") or "gemma4:latest"
OLLAMA_MODEL_UTILITY = env("OLLAMA_MODEL_UTILITY", "") or "llama3.2"
OLLAMA_EMBED_MODEL = env("OLLAMA_EMBED_MODEL", "") or "nomic-embed-text"
OLLAMA_NUM_CTX_THINKING = int(
env("OLLAMA_NUM_CTX_THINKING", "16384") or "16384"
)
OLLAMA_NUM_CTX_FAST = int(env("OLLAMA_NUM_CTX_FAST", "8192") or "8192")
OLLAMA_NUM_CTX_UTILITY = int(env("OLLAMA_NUM_CTX_UTILITY", "4096") or "4096")
CHROMA_PERSIST_DIRECTORY = env(
"CHROMA_PERSIST_DIRECTORY",
@@ -307,6 +316,15 @@ os.makedirs(directory_path, exist_ok=True)
ALLOW_IMAGE_GENERATION = env_bool("ALLOW_IMAGE_GENERATION", False)
ALLOW_INTERNET_ACCESS = env_bool("ALLOW_INTERNET_ACCESS", True)
# Web search (#62). Primary provider is SearxNG; DDGS is the automatic failover.
SEARCH_PROVIDER = (env("SEARCH_PROVIDER", "searxng") or "searxng").lower()
SEARCH_FAILOVER_PROVIDER = (
env("SEARCH_FAILOVER_PROVIDER", "ddgs") or "ddgs"
).lower()
SEARXNG_BASE_URL = (
env("SEARXNG_BASE_URL", "http://127.0.0.1:8080") or "http://127.0.0.1:8080"
).rstrip("/")
SEARXNG_TIMEOUT_SECONDS = float(env("SEARXNG_TIMEOUT_SECONDS", "8") or "8")
# When True, chat turns require an active plan and respect prompt/token quotas.
ENFORCE_SUBSCRIPTION_GATES = env_bool("ENFORCE_SUBSCRIPTION_GATES", True)
+1
View File
@@ -42,6 +42,7 @@ dependencies = [
"pillow==12.0.0",
"beautifulsoup4==4.14.3",
"ddgs==9.9.3",
"requests>=2.32,<3",
"httpx==0.28.1",
"python-dateutil==2.9.0.post0",
"pytz==2025.2",