Spaces:
Sleeping
Sleeping
Vineetiitg commited on
Commit ·
d816f3a
1
Parent(s): a30ab12
feat(backend): integrate Redis workers, persistent model caching, Cohere V2 fallback, and 5x TTFT fast-path RAG
Browse files- .env.example +12 -4
- .gitignore +2 -0
- Dockerfile.backend +0 -3
- README.md +1 -0
- app/auth/security.py +3 -3
- app/core/config.py +14 -7
- app/core/dependencies.py +3 -2
- app/engine/indexer.py +8 -2
- app/engine/ingestion.py +14 -1
- app/engine/memory.py +124 -2
- app/engine/query_transform.py +14 -8
- app/engine/reranker.py +86 -33
- app/engine/retriever.py +34 -1
- app/engine/semantic_cache.py +29 -7
- app/graph/workflow.py +53 -57
- app/guardrails/input.py +0 -1
- app/main.py +94 -16
- app/tests/eval_rag.py +2 -2
- requirements.txt +4 -2
.env.example
CHANGED
|
@@ -1,18 +1,26 @@
|
|
| 1 |
PROJECT_NAME="Support Docs Copilot"
|
| 2 |
OPENROUTER_API_KEY=your_openrouter_api_key_here
|
| 3 |
OPENROUTER_BASE_URL=https://openrouter.ai/api/v1
|
| 4 |
-
LLM_MODEL=
|
|
|
|
|
|
|
|
|
|
|
|
|
| 5 |
QDRANT_URL=
|
| 6 |
QDRANT_LOCATION=./qdrant_data
|
| 7 |
COLLECTION_NAME=support_docs
|
| 8 |
DATA_DIR=data/docs
|
|
|
|
|
|
|
|
|
|
|
|
|
| 9 |
RETRIEVAL_MODE=dense
|
| 10 |
-
RETRIEVAL_TOP_K=
|
| 11 |
RERANKER_TOP_N=3
|
| 12 |
-
RERANKER_ENABLED=
|
| 13 |
CHUNK_SIZE=500
|
| 14 |
CHUNK_OVERLAP=50
|
| 15 |
-
MIN_RELEVANCE_SCORE=0.
|
| 16 |
MAX_CONTEXT_CHARS=12000
|
| 17 |
ENABLE_GUARDRAILS=true
|
| 18 |
ENABLE_RAG_EVAL=false
|
|
|
|
| 1 |
PROJECT_NAME="Support Docs Copilot"
|
| 2 |
OPENROUTER_API_KEY=your_openrouter_api_key_here
|
| 3 |
OPENROUTER_BASE_URL=https://openrouter.ai/api/v1
|
| 4 |
+
LLM_MODEL=deepseek/deepseek-v4-flash
|
| 5 |
+
FAST_LLM_MODEL=google/gemini-2.0-flash-lite-preview-02-05:free
|
| 6 |
+
FAST_LLM_API_KEY=
|
| 7 |
+
FAST_LLM_BASE_URL=
|
| 8 |
+
SLOW_LLM_MODEL=deepseek/deepseek-r1
|
| 9 |
QDRANT_URL=
|
| 10 |
QDRANT_LOCATION=./qdrant_data
|
| 11 |
COLLECTION_NAME=support_docs
|
| 12 |
DATA_DIR=data/docs
|
| 13 |
+
COHERE_API_KEY=your_cohere_api_key_here
|
| 14 |
+
RERANKER_PROVIDER=auto
|
| 15 |
+
FLASHRANK_MODEL=ms-marco-TinyBERT-L-2-v2
|
| 16 |
+
RERANKER_MODEL=rerank-english-v3.0
|
| 17 |
RETRIEVAL_MODE=dense
|
| 18 |
+
RETRIEVAL_TOP_K=5
|
| 19 |
RERANKER_TOP_N=3
|
| 20 |
+
RERANKER_ENABLED=true
|
| 21 |
CHUNK_SIZE=500
|
| 22 |
CHUNK_OVERLAP=50
|
| 23 |
+
MIN_RELEVANCE_SCORE=0.2
|
| 24 |
MAX_CONTEXT_CHARS=12000
|
| 25 |
ENABLE_GUARDRAILS=true
|
| 26 |
ENABLE_RAG_EVAL=false
|
.gitignore
CHANGED
|
@@ -14,6 +14,8 @@ qdrant_data/
|
|
| 14 |
frontend.log
|
| 15 |
backend.log
|
| 16 |
fastembed_cache/
|
|
|
|
|
|
|
| 17 |
.cache/
|
| 18 |
reports/*.html
|
| 19 |
reports/*.json
|
|
|
|
| 14 |
frontend.log
|
| 15 |
backend.log
|
| 16 |
fastembed_cache/
|
| 17 |
+
*fastembed_cache/
|
| 18 |
+
*flashrank_cache/
|
| 19 |
.cache/
|
| 20 |
reports/*.html
|
| 21 |
reports/*.json
|
Dockerfile.backend
CHANGED
|
@@ -13,9 +13,6 @@ COPY ./tests /app/tests
|
|
| 13 |
COPY ./datasets /app/datasets
|
| 14 |
COPY ./data /app/data
|
| 15 |
|
| 16 |
-
RUN python -c "from langchain_community.embeddings import FastEmbedEmbeddings; FastEmbedEmbeddings(model_name='BAAI/bge-small-en-v1.5')"
|
| 17 |
-
RUN python -c "from sentence_transformers import CrossEncoder; CrossEncoder('cross-encoder/ms-marco-MiniLM-L-6-v2'); CrossEncoder('cross-encoder/nli-deberta-v3-small')"
|
| 18 |
-
|
| 19 |
EXPOSE 8000
|
| 20 |
|
| 21 |
CMD ["uvicorn", "app.main:app", "--host", "0.0.0.0", "--port", "8000"]
|
|
|
|
| 13 |
COPY ./datasets /app/datasets
|
| 14 |
COPY ./data /app/data
|
| 15 |
|
|
|
|
|
|
|
|
|
|
| 16 |
EXPOSE 8000
|
| 17 |
|
| 18 |
CMD ["uvicorn", "app.main:app", "--host", "0.0.0.0", "--port", "8000"]
|
README.md
CHANGED
|
@@ -200,3 +200,4 @@ make down # Tear down cluster and free ports
|
|
| 200 |
- **JWT Authentication:** Protected endpoints require OAuth2 Bearer Tokens. Authenticate via `/auth/login` (default roles: `user` and `admin`).
|
| 201 |
- **Input Guardrails:** Automatically checks for prompt injection and applies rate limiting (30 req/min).
|
| 202 |
- **Output Guardrails:** Automatically scrubs and redacts PII (SSNs, credit card numbers) before returning answers to the UI.
|
|
|
|
|
|
| 200 |
- **JWT Authentication:** Protected endpoints require OAuth2 Bearer Tokens. Authenticate via `/auth/login` (default roles: `user` and `admin`).
|
| 201 |
- **Input Guardrails:** Automatically checks for prompt injection and applies rate limiting (30 req/min).
|
| 202 |
- **Output Guardrails:** Automatically scrubs and redacts PII (SSNs, credit card numbers) before returning answers to the UI.
|
| 203 |
+
.
|
app/auth/security.py
CHANGED
|
@@ -33,17 +33,17 @@ def resolve_user(token: str | None = Depends(oauth2_scheme)) -> UserContext:
|
|
| 33 |
if not settings.AUTH_ENABLED:
|
| 34 |
return UserContext(role="admin", user_id="local-dev")
|
| 35 |
if not token:
|
| 36 |
-
|
| 37 |
|
| 38 |
try:
|
| 39 |
payload = jwt.decode(token, settings.SECRET_KEY, algorithms=["HS256"])
|
| 40 |
username: str | None = payload.get("sub")
|
| 41 |
role: str | None = payload.get("role")
|
| 42 |
if username is None or role is None:
|
| 43 |
-
|
| 44 |
return UserContext(role=role, user_id=username)
|
| 45 |
except jwt.PyJWTError:
|
| 46 |
-
|
| 47 |
|
| 48 |
|
| 49 |
def require_admin(user: UserContext = Depends(resolve_user)) -> None:
|
|
|
|
| 33 |
if not settings.AUTH_ENABLED:
|
| 34 |
return UserContext(role="admin", user_id="local-dev")
|
| 35 |
if not token:
|
| 36 |
+
return UserContext(role="guest", user_id="guest")
|
| 37 |
|
| 38 |
try:
|
| 39 |
payload = jwt.decode(token, settings.SECRET_KEY, algorithms=["HS256"])
|
| 40 |
username: str | None = payload.get("sub")
|
| 41 |
role: str | None = payload.get("role")
|
| 42 |
if username is None or role is None:
|
| 43 |
+
return UserContext(role="guest", user_id="guest")
|
| 44 |
return UserContext(role=role, user_id=username)
|
| 45 |
except jwt.PyJWTError:
|
| 46 |
+
return UserContext(role="guest", user_id="guest")
|
| 47 |
|
| 48 |
|
| 49 |
def require_admin(user: UserContext = Depends(resolve_user)) -> None:
|
app/core/config.py
CHANGED
|
@@ -3,11 +3,13 @@ from pydantic_settings import BaseSettings
|
|
| 3 |
class Settings(BaseSettings):
|
| 4 |
PROJECT_NAME: str = "Support Docs Copilot"
|
| 5 |
|
| 6 |
-
# OpenRouter LLM Config
|
| 7 |
OPENROUTER_API_KEY: str = ""
|
| 8 |
OPENROUTER_BASE_URL: str = "https://openrouter.ai/api/v1"
|
| 9 |
LLM_MODEL: str = "deepseek/deepseek-v4-flash"
|
| 10 |
-
FAST_LLM_MODEL: str = "
|
|
|
|
|
|
|
| 11 |
SLOW_LLM_MODEL: str = "deepseek/deepseek-r1"
|
| 12 |
|
| 13 |
# Qdrant Vector DB Config
|
|
@@ -19,17 +21,22 @@ class Settings(BaseSettings):
|
|
| 19 |
# Redis & Queue Config
|
| 20 |
REDIS_URL: str = "redis://redis:6379/0"
|
| 21 |
|
| 22 |
-
#
|
|
|
|
|
|
|
|
|
|
| 23 |
DENSE_EMBEDDING_MODEL: str = "BAAI/bge-small-en-v1.5"
|
| 24 |
SPARSE_EMBEDDING_MODEL: str = "Qdrant/bm25"
|
| 25 |
-
|
|
|
|
|
|
|
| 26 |
|
| 27 |
# Retrieval Config
|
| 28 |
RETRIEVAL_MODE: str = "dense"
|
| 29 |
-
RETRIEVAL_TOP_K: int =
|
| 30 |
RERANKER_TOP_N: int = 3
|
| 31 |
-
RERANKER_ENABLED: bool =
|
| 32 |
-
MIN_RELEVANCE_SCORE: float = 0.
|
| 33 |
MAX_CONTEXT_CHARS: int = 12000
|
| 34 |
|
| 35 |
# Chunking Config
|
|
|
|
| 3 |
class Settings(BaseSettings):
|
| 4 |
PROJECT_NAME: str = "Support Docs Copilot"
|
| 5 |
|
| 6 |
+
# OpenRouter / LPU LLM Config
|
| 7 |
OPENROUTER_API_KEY: str = ""
|
| 8 |
OPENROUTER_BASE_URL: str = "https://openrouter.ai/api/v1"
|
| 9 |
LLM_MODEL: str = "deepseek/deepseek-v4-flash"
|
| 10 |
+
FAST_LLM_MODEL: str = "google/gemini-2.0-flash-lite-preview-02-05"
|
| 11 |
+
FAST_LLM_API_KEY: str = ""
|
| 12 |
+
FAST_LLM_BASE_URL: str = ""
|
| 13 |
SLOW_LLM_MODEL: str = "deepseek/deepseek-r1"
|
| 14 |
|
| 15 |
# Qdrant Vector DB Config
|
|
|
|
| 21 |
# Redis & Queue Config
|
| 22 |
REDIS_URL: str = "redis://redis:6379/0"
|
| 23 |
|
| 24 |
+
# Cohere API Config (for Document Ranking and Relevance Grading)
|
| 25 |
+
COHERE_API_KEY: str = ""
|
| 26 |
+
|
| 27 |
+
# Embeddings & Reranker Config
|
| 28 |
DENSE_EMBEDDING_MODEL: str = "BAAI/bge-small-en-v1.5"
|
| 29 |
SPARSE_EMBEDDING_MODEL: str = "Qdrant/bm25"
|
| 30 |
+
RERANKER_PROVIDER: str = "auto" # "auto" (cohere if key present else flashrank), "cohere", "flashrank"
|
| 31 |
+
RERANKER_MODEL: str = "rerank-english-v3.0"
|
| 32 |
+
FLASHRANK_MODEL: str = "ms-marco-TinyBERT-L-2-v2"
|
| 33 |
|
| 34 |
# Retrieval Config
|
| 35 |
RETRIEVAL_MODE: str = "dense"
|
| 36 |
+
RETRIEVAL_TOP_K: int = 5
|
| 37 |
RERANKER_TOP_N: int = 3
|
| 38 |
+
RERANKER_ENABLED: bool = True
|
| 39 |
+
MIN_RELEVANCE_SCORE: float = 0.2
|
| 40 |
MAX_CONTEXT_CHARS: int = 12000
|
| 41 |
|
| 42 |
# Chunking Config
|
app/core/dependencies.py
CHANGED
|
@@ -24,8 +24,9 @@ def get_qdrant_client() -> QdrantClient:
|
|
| 24 |
def check_openrouter() -> dict[str, Any]:
|
| 25 |
try:
|
| 26 |
headers = {"Authorization": f"Bearer {settings.OPENROUTER_API_KEY}"}
|
| 27 |
-
|
| 28 |
-
|
|
|
|
| 29 |
except Exception as exc:
|
| 30 |
return {"ok": bool(settings.OPENROUTER_API_KEY), "error": str(exc)}
|
| 31 |
|
|
|
|
| 24 |
def check_openrouter() -> dict[str, Any]:
|
| 25 |
try:
|
| 26 |
headers = {"Authorization": f"Bearer {settings.OPENROUTER_API_KEY}"}
|
| 27 |
+
url = f"{settings.OPENROUTER_BASE_URL.rstrip('/v1').rstrip('/')}/api/v1/auth/key" if "openrouter.ai" in settings.OPENROUTER_BASE_URL else f"{settings.OPENROUTER_BASE_URL}/models"
|
| 28 |
+
response = requests.get(url, headers=headers, timeout=5)
|
| 29 |
+
return {"ok": response.ok or response.status_code == 200, "status_code": response.status_code}
|
| 30 |
except Exception as exc:
|
| 31 |
return {"ok": bool(settings.OPENROUTER_API_KEY), "error": str(exc)}
|
| 32 |
|
app/engine/indexer.py
CHANGED
|
@@ -21,14 +21,20 @@ _sparse_embedder = None
|
|
| 21 |
def dense_embeddings() -> FastEmbedEmbeddings:
|
| 22 |
global _dense_embedder
|
| 23 |
if _dense_embedder is None:
|
| 24 |
-
|
|
|
|
|
|
|
|
|
|
| 25 |
return _dense_embedder
|
| 26 |
|
| 27 |
|
| 28 |
def sparse_embeddings() -> FastEmbedSparse:
|
| 29 |
global _sparse_embedder
|
| 30 |
if _sparse_embedder is None:
|
| 31 |
-
|
|
|
|
|
|
|
|
|
|
| 32 |
return _sparse_embedder
|
| 33 |
|
| 34 |
|
|
|
|
| 21 |
def dense_embeddings() -> FastEmbedEmbeddings:
|
| 22 |
global _dense_embedder
|
| 23 |
if _dense_embedder is None:
|
| 24 |
+
import os
|
| 25 |
+
cache_dir = "/app/data/fastembed_cache" if os.path.exists("/app") else "./data/fastembed_cache"
|
| 26 |
+
os.makedirs(cache_dir, exist_ok=True)
|
| 27 |
+
_dense_embedder = FastEmbedEmbeddings(model_name=settings.DENSE_EMBEDDING_MODEL, cache_dir=cache_dir)
|
| 28 |
return _dense_embedder
|
| 29 |
|
| 30 |
|
| 31 |
def sparse_embeddings() -> FastEmbedSparse:
|
| 32 |
global _sparse_embedder
|
| 33 |
if _sparse_embedder is None:
|
| 34 |
+
import os
|
| 35 |
+
cache_dir = "/app/data/fastembed_cache" if os.path.exists("/app") else "./data/fastembed_cache"
|
| 36 |
+
os.makedirs(cache_dir, exist_ok=True)
|
| 37 |
+
_sparse_embedder = FastEmbedSparse(model_name=settings.SPARSE_EMBEDDING_MODEL, cache_dir=cache_dir)
|
| 38 |
return _sparse_embedder
|
| 39 |
|
| 40 |
|
app/engine/ingestion.py
CHANGED
|
@@ -92,10 +92,23 @@ def delete_indexed_document(doc_id: str) -> None:
|
|
| 92 |
if doc_id not in registry:
|
| 93 |
logger.info(f"Document not found: {doc_id}")
|
| 94 |
return
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 95 |
delete_document(doc_id)
|
| 96 |
del registry[doc_id]
|
| 97 |
save_registry(registry)
|
| 98 |
-
logger.info(f"Deleted document: {doc_id}")
|
| 99 |
|
| 100 |
|
| 101 |
def reset_index() -> None:
|
|
|
|
| 92 |
if doc_id not in registry:
|
| 93 |
logger.info(f"Document not found: {doc_id}")
|
| 94 |
return
|
| 95 |
+
record = registry[doc_id]
|
| 96 |
+
if source_path := record.get("source_path"):
|
| 97 |
+
try:
|
| 98 |
+
Path(source_path).unlink(missing_ok=True)
|
| 99 |
+
logger.info(f"Deleted physical file: {source_path}")
|
| 100 |
+
except Exception as e:
|
| 101 |
+
logger.warning(f"Failed to delete file {source_path}: {e}")
|
| 102 |
+
elif source := record.get("source"):
|
| 103 |
+
try:
|
| 104 |
+
(Path("data/docs") / source).unlink(missing_ok=True)
|
| 105 |
+
logger.info(f"Deleted physical file from data/docs/: {source}")
|
| 106 |
+
except Exception as e:
|
| 107 |
+
logger.warning(f"Failed to delete file {source}: {e}")
|
| 108 |
delete_document(doc_id)
|
| 109 |
del registry[doc_id]
|
| 110 |
save_registry(registry)
|
| 111 |
+
logger.info(f"Deleted document and vector embeddings: {doc_id}")
|
| 112 |
|
| 113 |
|
| 114 |
def reset_index() -> None:
|
app/engine/memory.py
CHANGED
|
@@ -1,8 +1,12 @@
|
|
|
|
|
| 1 |
import json
|
| 2 |
import logging
|
| 3 |
import time
|
| 4 |
import uuid
|
| 5 |
from typing import Any, Dict, List, Optional
|
|
|
|
|
|
|
|
|
|
| 6 |
from app.core.queue import get_redis_client
|
| 7 |
|
| 8 |
logger = logging.getLogger(__name__)
|
|
@@ -42,14 +46,96 @@ async def add_session_message(
|
|
| 42 |
}
|
| 43 |
await redis.setex(meta_key, SESSION_TTL, json.dumps(meta))
|
| 44 |
|
| 45 |
-
# Add to user's list of sessions
|
| 46 |
await redis.zadd(user_sessions_key, {session_id: time.time()})
|
| 47 |
await redis.expire(user_sessions_key, SESSION_TTL)
|
|
|
|
|
|
|
| 48 |
logger.debug(f"Added message to session {session_id} for user {user_id}")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 49 |
except Exception as e:
|
| 50 |
logger.warning(f"Failed to add session message to Redis: {e}")
|
| 51 |
|
| 52 |
-
async def
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 53 |
try:
|
| 54 |
redis = await get_redis_client()
|
| 55 |
key = f"session:{user_id}:{session_id}:messages"
|
|
@@ -93,16 +179,52 @@ async def list_user_sessions(user_id: str, limit: int = 20) -> List[Dict[str, An
|
|
| 93 |
logger.warning(f"Failed to list user sessions from Redis: {e}")
|
| 94 |
return []
|
| 95 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 96 |
async def delete_session(user_id: str, session_id: str) -> bool:
|
| 97 |
try:
|
| 98 |
redis = await get_redis_client()
|
| 99 |
key = f"session:{user_id}:{session_id}:messages"
|
| 100 |
meta_key = f"session:{user_id}:{session_id}:meta"
|
|
|
|
| 101 |
user_sessions_key = f"user_sessions:{user_id}"
|
| 102 |
|
| 103 |
await redis.delete(key)
|
| 104 |
await redis.delete(meta_key)
|
|
|
|
| 105 |
await redis.zrem(user_sessions_key, session_id)
|
|
|
|
| 106 |
logger.info(f"Deleted session {session_id} for user {user_id}")
|
| 107 |
return True
|
| 108 |
except Exception as e:
|
|
|
|
| 1 |
+
import asyncio
|
| 2 |
import json
|
| 3 |
import logging
|
| 4 |
import time
|
| 5 |
import uuid
|
| 6 |
from typing import Any, Dict, List, Optional
|
| 7 |
+
from langchain_openai import ChatOpenAI
|
| 8 |
+
from langchain_core.prompts import PromptTemplate
|
| 9 |
+
from app.core.config import settings
|
| 10 |
from app.core.queue import get_redis_client
|
| 11 |
|
| 12 |
logger = logging.getLogger(__name__)
|
|
|
|
| 46 |
}
|
| 47 |
await redis.setex(meta_key, SESSION_TTL, json.dumps(meta))
|
| 48 |
|
| 49 |
+
# Add to user's list of sessions and global admin list
|
| 50 |
await redis.zadd(user_sessions_key, {session_id: time.time()})
|
| 51 |
await redis.expire(user_sessions_key, SESSION_TTL)
|
| 52 |
+
await redis.zadd("all_sessions", {f"{user_id}:{session_id}": time.time()})
|
| 53 |
+
await redis.expire("all_sessions", SESSION_TTL)
|
| 54 |
logger.debug(f"Added message to session {session_id} for user {user_id}")
|
| 55 |
+
|
| 56 |
+
# Trigger background summarization if chat exceeds 6 turns (sliding window)
|
| 57 |
+
total_msgs = await redis.llen(key)
|
| 58 |
+
if total_msgs > 6:
|
| 59 |
+
asyncio.create_task(summarize_session_if_needed(user_id, session_id))
|
| 60 |
except Exception as e:
|
| 61 |
logger.warning(f"Failed to add session message to Redis: {e}")
|
| 62 |
|
| 63 |
+
async def get_session_summary(user_id: str, session_id: str) -> Optional[str]:
|
| 64 |
+
try:
|
| 65 |
+
redis = await get_redis_client()
|
| 66 |
+
summary_key = f"session:{user_id}:{session_id}:summary"
|
| 67 |
+
raw_summary = await redis.get(summary_key)
|
| 68 |
+
if raw_summary:
|
| 69 |
+
return raw_summary.decode("utf-8") if isinstance(raw_summary, bytes) else str(raw_summary)
|
| 70 |
+
return None
|
| 71 |
+
except Exception as e:
|
| 72 |
+
logger.warning(f"Failed to get session summary from Redis: {e}")
|
| 73 |
+
return None
|
| 74 |
+
|
| 75 |
+
async def summarize_session_if_needed(user_id: str, session_id: str) -> None:
|
| 76 |
+
try:
|
| 77 |
+
redis = await get_redis_client()
|
| 78 |
+
key = f"session:{user_id}:{session_id}:messages"
|
| 79 |
+
summary_key = f"session:{user_id}:{session_id}:summary"
|
| 80 |
+
|
| 81 |
+
total_len = await redis.llen(key)
|
| 82 |
+
if total_len <= 6:
|
| 83 |
+
return
|
| 84 |
+
|
| 85 |
+
# Get older turns (turns 1 through N-6)
|
| 86 |
+
older_raw = await redis.lrange(key, 0, -7)
|
| 87 |
+
if not older_raw:
|
| 88 |
+
return
|
| 89 |
+
|
| 90 |
+
older_msgs = []
|
| 91 |
+
for raw in older_raw:
|
| 92 |
+
try:
|
| 93 |
+
older_msgs.append(json.loads(raw))
|
| 94 |
+
except Exception:
|
| 95 |
+
continue
|
| 96 |
+
|
| 97 |
+
if not older_msgs:
|
| 98 |
+
return
|
| 99 |
+
|
| 100 |
+
old_summary = await get_session_summary(user_id, session_id)
|
| 101 |
+
lines = []
|
| 102 |
+
if old_summary:
|
| 103 |
+
lines.append(f"Previous Summary: {old_summary}")
|
| 104 |
+
for m in older_msgs:
|
| 105 |
+
lines.append(f"{m.get('role', 'user')}: {m.get('content', '')}")
|
| 106 |
+
|
| 107 |
+
to_summarize = "\n".join(lines)
|
| 108 |
+
|
| 109 |
+
llm = ChatOpenAI(
|
| 110 |
+
model=getattr(settings, "FAST_LLM_MODEL", settings.LLM_MODEL),
|
| 111 |
+
temperature=0,
|
| 112 |
+
openai_api_key=getattr(settings, "FAST_LLM_API_KEY", "") or settings.OPENROUTER_API_KEY,
|
| 113 |
+
openai_api_base=getattr(settings, "FAST_LLM_BASE_URL", "") or settings.OPENROUTER_BASE_URL,
|
| 114 |
+
default_headers={"HTTP-Referer": "https://localhost:3000", "X-Title": "Support Docs Copilot"},
|
| 115 |
+
)
|
| 116 |
+
prompt = PromptTemplate(
|
| 117 |
+
template="""You are a helpful technical conversation summarizer.
|
| 118 |
+
Condense the following past conversation turns into a dense, 2-sentence summary capturing the main technical topics, user goals, and key details discussed.
|
| 119 |
+
Your output must start exactly with: "System Summary:"
|
| 120 |
+
|
| 121 |
+
Conversation to summarize:
|
| 122 |
+
{text_to_summarize}
|
| 123 |
+
|
| 124 |
+
Summary:""",
|
| 125 |
+
input_variables=["text_to_summarize"],
|
| 126 |
+
)
|
| 127 |
+
chain = prompt | llm
|
| 128 |
+
res = await chain.ainvoke({"text_to_summarize": to_summarize})
|
| 129 |
+
summary_text = res.content.strip()
|
| 130 |
+
if not summary_text.startswith("System Summary:"):
|
| 131 |
+
summary_text = f"System Summary: {summary_text}"
|
| 132 |
+
|
| 133 |
+
await redis.setex(summary_key, SESSION_TTL, summary_text)
|
| 134 |
+
logger.info(f"Generated background conversation summary for session {session_id}")
|
| 135 |
+
except Exception as e:
|
| 136 |
+
logger.warning(f"Failed to summarize session history in background: {e}")
|
| 137 |
+
|
| 138 |
+
async def get_session_history(user_id: str, session_id: str, limit: int = 6) -> List[Dict[str, Any]]:
|
| 139 |
try:
|
| 140 |
redis = await get_redis_client()
|
| 141 |
key = f"session:{user_id}:{session_id}:messages"
|
|
|
|
| 179 |
logger.warning(f"Failed to list user sessions from Redis: {e}")
|
| 180 |
return []
|
| 181 |
|
| 182 |
+
async def list_all_sessions(limit: int = 50) -> List[Dict[str, Any]]:
|
| 183 |
+
try:
|
| 184 |
+
redis = await get_redis_client()
|
| 185 |
+
entries = await redis.zrevrange("all_sessions", 0, limit - 1)
|
| 186 |
+
if not entries:
|
| 187 |
+
keys = await redis.keys("user_sessions:*")
|
| 188 |
+
sessions = []
|
| 189 |
+
for k in keys:
|
| 190 |
+
uid = k.decode("utf-8").split(":")[-1] if isinstance(k, bytes) else str(k).split(":")[-1]
|
| 191 |
+
user_sess = await list_user_sessions(uid, limit=10)
|
| 192 |
+
sessions.extend(user_sess)
|
| 193 |
+
sessions.sort(key=lambda x: x.get("updated_at", 0), reverse=True)
|
| 194 |
+
return sessions[:limit]
|
| 195 |
+
|
| 196 |
+
sessions = []
|
| 197 |
+
for entry_bytes in entries:
|
| 198 |
+
entry_str = entry_bytes.decode("utf-8") if isinstance(entry_bytes, bytes) else str(entry_bytes)
|
| 199 |
+
if ":" in entry_str:
|
| 200 |
+
uid, sid = entry_str.split(":", 1)
|
| 201 |
+
meta_key = f"session:{uid}:{sid}:meta"
|
| 202 |
+
raw_meta = await redis.get(meta_key)
|
| 203 |
+
if raw_meta:
|
| 204 |
+
try:
|
| 205 |
+
sessions.append(json.loads(raw_meta))
|
| 206 |
+
except Exception:
|
| 207 |
+
sessions.append({"session_id": sid, "user_id": uid, "updated_at": time.time(), "last_preview": "Chat Session"})
|
| 208 |
+
else:
|
| 209 |
+
sessions.append({"session_id": sid, "user_id": uid, "updated_at": time.time(), "last_preview": "Chat Session"})
|
| 210 |
+
return sessions
|
| 211 |
+
except Exception as e:
|
| 212 |
+
logger.warning(f"Failed to list all sessions from Redis: {e}")
|
| 213 |
+
return []
|
| 214 |
+
|
| 215 |
async def delete_session(user_id: str, session_id: str) -> bool:
|
| 216 |
try:
|
| 217 |
redis = await get_redis_client()
|
| 218 |
key = f"session:{user_id}:{session_id}:messages"
|
| 219 |
meta_key = f"session:{user_id}:{session_id}:meta"
|
| 220 |
+
summary_key = f"session:{user_id}:{session_id}:summary"
|
| 221 |
user_sessions_key = f"user_sessions:{user_id}"
|
| 222 |
|
| 223 |
await redis.delete(key)
|
| 224 |
await redis.delete(meta_key)
|
| 225 |
+
await redis.delete(summary_key)
|
| 226 |
await redis.zrem(user_sessions_key, session_id)
|
| 227 |
+
await redis.zrem("all_sessions", f"{user_id}:{session_id}")
|
| 228 |
logger.info(f"Deleted session {session_id} for user {user_id}")
|
| 229 |
return True
|
| 230 |
except Exception as e:
|
app/engine/query_transform.py
CHANGED
|
@@ -23,8 +23,8 @@ async def query_variants(query: str, chat_history: list[dict] = None) -> list[st
|
|
| 23 |
llm = ChatOpenAI(
|
| 24 |
model=getattr(settings, "FAST_LLM_MODEL", settings.LLM_MODEL),
|
| 25 |
temperature=0,
|
| 26 |
-
openai_api_key=settings.OPENROUTER_API_KEY,
|
| 27 |
-
openai_api_base=settings.OPENROUTER_BASE_URL,
|
| 28 |
default_headers={"HTTP-Referer": "https://localhost:3000", "X-Title": "Support Docs Copilot"},
|
| 29 |
)
|
| 30 |
prompt = PromptTemplate(
|
|
@@ -41,7 +41,10 @@ User Question: {question}""",
|
|
| 41 |
chain = prompt | llm
|
| 42 |
result = await chain.ainvoke({"question": normalized, "chat_history": history_str})
|
| 43 |
|
| 44 |
-
|
|
|
|
|
|
|
|
|
|
| 45 |
new_variants = parsed.get("variants", [])
|
| 46 |
|
| 47 |
if isinstance(new_variants, list):
|
|
@@ -55,19 +58,22 @@ User Question: {question}""",
|
|
| 55 |
return list(dict.fromkeys(variants))
|
| 56 |
|
| 57 |
|
| 58 |
-
async def condense_query(query: str, chat_history: list[dict] = None) -> str:
|
| 59 |
-
if not chat_history:
|
| 60 |
return normalize_query(query)
|
| 61 |
|
| 62 |
normalized = normalize_query(query)
|
| 63 |
-
|
|
|
|
|
|
|
|
|
|
| 64 |
|
| 65 |
try:
|
| 66 |
llm = ChatOpenAI(
|
| 67 |
model=getattr(settings, "FAST_LLM_MODEL", settings.LLM_MODEL),
|
| 68 |
temperature=0,
|
| 69 |
-
openai_api_key=settings.OPENROUTER_API_KEY,
|
| 70 |
-
openai_api_base=settings.OPENROUTER_BASE_URL,
|
| 71 |
default_headers={"HTTP-Referer": "https://localhost:3000", "X-Title": "Support Docs Copilot"},
|
| 72 |
)
|
| 73 |
prompt = PromptTemplate(
|
|
|
|
| 23 |
llm = ChatOpenAI(
|
| 24 |
model=getattr(settings, "FAST_LLM_MODEL", settings.LLM_MODEL),
|
| 25 |
temperature=0,
|
| 26 |
+
openai_api_key=getattr(settings, "FAST_LLM_API_KEY", "") or settings.OPENROUTER_API_KEY,
|
| 27 |
+
openai_api_base=getattr(settings, "FAST_LLM_BASE_URL", "") or settings.OPENROUTER_BASE_URL,
|
| 28 |
default_headers={"HTTP-Referer": "https://localhost:3000", "X-Title": "Support Docs Copilot"},
|
| 29 |
)
|
| 30 |
prompt = PromptTemplate(
|
|
|
|
| 41 |
chain = prompt | llm
|
| 42 |
result = await chain.ainvoke({"question": normalized, "chat_history": history_str})
|
| 43 |
|
| 44 |
+
content = result.content.strip()
|
| 45 |
+
if content.startswith("```"):
|
| 46 |
+
content = re.sub(r"^```(?:json)?\s*|\s*```$", "", content, flags=re.MULTILINE).strip()
|
| 47 |
+
parsed = json.loads(content)
|
| 48 |
new_variants = parsed.get("variants", [])
|
| 49 |
|
| 50 |
if isinstance(new_variants, list):
|
|
|
|
| 58 |
return list(dict.fromkeys(variants))
|
| 59 |
|
| 60 |
|
| 61 |
+
async def condense_query(query: str, chat_history: list[dict] = None, summary: str = None) -> str:
|
| 62 |
+
if not chat_history and not summary:
|
| 63 |
return normalize_query(query)
|
| 64 |
|
| 65 |
normalized = normalize_query(query)
|
| 66 |
+
lines = [f"{msg['role']}: {msg['content']}" for msg in (chat_history or [])[-6:]]
|
| 67 |
+
if summary:
|
| 68 |
+
lines.insert(0, summary if summary.startswith("System Summary:") else f"System Summary: {summary}")
|
| 69 |
+
history_str = "\n".join(lines)
|
| 70 |
|
| 71 |
try:
|
| 72 |
llm = ChatOpenAI(
|
| 73 |
model=getattr(settings, "FAST_LLM_MODEL", settings.LLM_MODEL),
|
| 74 |
temperature=0,
|
| 75 |
+
openai_api_key=getattr(settings, "FAST_LLM_API_KEY", "") or settings.OPENROUTER_API_KEY,
|
| 76 |
+
openai_api_base=getattr(settings, "FAST_LLM_BASE_URL", "") or settings.OPENROUTER_BASE_URL,
|
| 77 |
default_headers={"HTTP-Referer": "https://localhost:3000", "X-Title": "Support Docs Copilot"},
|
| 78 |
)
|
| 79 |
prompt = PromptTemplate(
|
app/engine/reranker.py
CHANGED
|
@@ -1,63 +1,116 @@
|
|
| 1 |
import logging
|
| 2 |
from typing import List, Tuple
|
| 3 |
from langchain_core.documents import Document
|
|
|
|
| 4 |
|
| 5 |
logger = logging.getLogger(__name__)
|
| 6 |
|
| 7 |
-
|
|
|
|
| 8 |
_nli_model = None
|
| 9 |
|
| 10 |
|
| 11 |
-
def
|
| 12 |
-
global
|
| 13 |
-
if
|
|
|
|
|
|
|
|
|
|
| 14 |
try:
|
| 15 |
-
|
| 16 |
-
logger.info("
|
| 17 |
-
|
| 18 |
except Exception as e:
|
| 19 |
-
logger.error(f"Failed to
|
| 20 |
-
|
| 21 |
-
return
|
| 22 |
|
| 23 |
|
| 24 |
-
def
|
| 25 |
-
global
|
| 26 |
-
if
|
| 27 |
try:
|
| 28 |
-
from
|
| 29 |
-
|
| 30 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 31 |
except Exception as e:
|
| 32 |
-
logger.error(f"Failed to
|
| 33 |
-
|
| 34 |
-
return
|
| 35 |
|
| 36 |
|
| 37 |
-
def
|
| 38 |
-
|
| 39 |
-
|
| 40 |
-
|
|
|
|
| 41 |
try:
|
| 42 |
-
|
| 43 |
-
|
| 44 |
-
|
| 45 |
-
|
| 46 |
-
|
| 47 |
-
|
|
|
|
| 48 |
|
| 49 |
reranked = []
|
| 50 |
-
for
|
| 51 |
-
|
|
|
|
|
|
|
|
|
|
| 52 |
reranked.append(doc)
|
| 53 |
|
| 54 |
-
|
|
|
|
| 55 |
return reranked
|
| 56 |
except Exception as e:
|
| 57 |
-
logger.
|
| 58 |
return documents[:top_k]
|
| 59 |
|
| 60 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 61 |
def evaluate_nli_groundedness(premise: str, hypothesis: str) -> Tuple[str, float]:
|
| 62 |
try:
|
| 63 |
model = get_nli_model()
|
|
|
|
| 1 |
import logging
|
| 2 |
from typing import List, Tuple
|
| 3 |
from langchain_core.documents import Document
|
| 4 |
+
from app.core.config import settings
|
| 5 |
|
| 6 |
logger = logging.getLogger(__name__)
|
| 7 |
|
| 8 |
+
_cohere_client = None
|
| 9 |
+
_flashrank_client = None
|
| 10 |
_nli_model = None
|
| 11 |
|
| 12 |
|
| 13 |
+
def get_cohere_client():
|
| 14 |
+
global _cohere_client
|
| 15 |
+
if _cohere_client is None:
|
| 16 |
+
if not settings.COHERE_API_KEY or "your_cohere" in settings.COHERE_API_KEY:
|
| 17 |
+
logger.warning("COHERE_API_KEY not set or default! Cohere reranking disabled.")
|
| 18 |
+
return None
|
| 19 |
try:
|
| 20 |
+
import cohere
|
| 21 |
+
logger.info(f"Initializing Cohere API Client with model: {settings.RERANKER_MODEL}")
|
| 22 |
+
_cohere_client = cohere.ClientV2(api_key=settings.COHERE_API_KEY)
|
| 23 |
except Exception as e:
|
| 24 |
+
logger.error(f"Failed to initialize Cohere client: {e}")
|
| 25 |
+
return None
|
| 26 |
+
return _cohere_client
|
| 27 |
|
| 28 |
|
| 29 |
+
def get_flashrank_client():
|
| 30 |
+
global _flashrank_client
|
| 31 |
+
if _flashrank_client is None:
|
| 32 |
try:
|
| 33 |
+
from flashrank import Ranker
|
| 34 |
+
import os
|
| 35 |
+
cache_dir = os.environ.get("FLASHRANK_CACHE_DIR", "/app/data/flashrank_cache" if os.path.exists("/app") else "./data/flashrank_cache")
|
| 36 |
+
os.makedirs(cache_dir, exist_ok=True)
|
| 37 |
+
model_name = getattr(settings, "FLASHRANK_MODEL", "ms-marco-TinyBERT-L-2-v2")
|
| 38 |
+
logger.info(f"Initializing local FlashRank ONNX Client with model: {model_name}")
|
| 39 |
+
_flashrank_client = Ranker(model_name=model_name, cache_dir=cache_dir)
|
| 40 |
except Exception as e:
|
| 41 |
+
logger.error(f"Failed to initialize FlashRank client: {e}")
|
| 42 |
+
return None
|
| 43 |
+
return _flashrank_client
|
| 44 |
|
| 45 |
|
| 46 |
+
def rerank_with_flashrank(question: str, documents: List[Document], top_k: int) -> List[Document]:
|
| 47 |
+
client = get_flashrank_client()
|
| 48 |
+
if not client:
|
| 49 |
+
logger.warning("FlashRank unavailable, returning un-reranked documents.")
|
| 50 |
+
return documents[:top_k]
|
| 51 |
try:
|
| 52 |
+
from flashrank import RerankRequest
|
| 53 |
+
passages = [
|
| 54 |
+
{"id": str(i), "text": doc.page_content, "meta": doc.metadata}
|
| 55 |
+
for i, doc in enumerate(documents)
|
| 56 |
+
]
|
| 57 |
+
request = RerankRequest(query=question, passages=passages)
|
| 58 |
+
results = client.rerank(request)[:min(top_k, len(documents))]
|
| 59 |
|
| 60 |
reranked = []
|
| 61 |
+
for r in results:
|
| 62 |
+
idx = int(r["id"])
|
| 63 |
+
doc = documents[idx]
|
| 64 |
+
doc.metadata["rerank_score"] = float(r["score"])
|
| 65 |
+
doc.metadata["relevance_score"] = float(r["score"])
|
| 66 |
reranked.append(doc)
|
| 67 |
|
| 68 |
+
highest_score = reranked[0].metadata["rerank_score"] if reranked else 0.0
|
| 69 |
+
logger.info(f"FlashRank ONNX Reranked {len(documents)} docs down to top-{len(reranked)} (highest score: {highest_score:.4f})")
|
| 70 |
return reranked
|
| 71 |
except Exception as e:
|
| 72 |
+
logger.error(f"FlashRank reranking failed ({e}), falling back to top_k truncation without scoring.")
|
| 73 |
return documents[:top_k]
|
| 74 |
|
| 75 |
|
| 76 |
+
def rerank_documents(question: str, documents: List[Document], top_k: int = 3) -> List[Document]:
|
| 77 |
+
if not documents or len(documents) <= 1:
|
| 78 |
+
return documents
|
| 79 |
+
|
| 80 |
+
provider = getattr(settings, "RERANKER_PROVIDER", "auto").lower()
|
| 81 |
+
has_cohere_key = bool(settings.COHERE_API_KEY and settings.COHERE_API_KEY.strip() != "" and "your_cohere" not in settings.COHERE_API_KEY)
|
| 82 |
+
|
| 83 |
+
# 1. Try Cohere API if explicitly selected or if 'auto' with a valid API key
|
| 84 |
+
if provider == "cohere" or (provider == "auto" and has_cohere_key):
|
| 85 |
+
client = get_cohere_client()
|
| 86 |
+
if client:
|
| 87 |
+
try:
|
| 88 |
+
doc_texts = [doc.page_content for doc in documents]
|
| 89 |
+
response = client.rerank(
|
| 90 |
+
model=settings.RERANKER_MODEL,
|
| 91 |
+
query=question,
|
| 92 |
+
documents=doc_texts,
|
| 93 |
+
top_n=min(top_k, len(documents))
|
| 94 |
+
)
|
| 95 |
+
|
| 96 |
+
reranked = []
|
| 97 |
+
for r in response.results:
|
| 98 |
+
doc = documents[r.index]
|
| 99 |
+
doc.metadata["rerank_score"] = float(r.relevance_score)
|
| 100 |
+
doc.metadata["relevance_score"] = float(r.relevance_score)
|
| 101 |
+
reranked.append(doc)
|
| 102 |
+
|
| 103 |
+
highest_score = reranked[0].metadata["rerank_score"] if reranked else 0.0
|
| 104 |
+
logger.info(f"Cohere API Reranked {len(documents)} docs down to top-{len(reranked)} (highest score: {highest_score:.4f})")
|
| 105 |
+
return reranked
|
| 106 |
+
except Exception as e:
|
| 107 |
+
logger.warning(f"Cohere API reranking failed ({e}). Attempting seamless fallback to local FlashRank...")
|
| 108 |
+
|
| 109 |
+
# 2. Use local FlashRank ONNX reranker (if provider=='flashrank', no Cohere key, or Cohere API fallback)
|
| 110 |
+
logger.info("Using local FlashRank CPU ONNX reranker.")
|
| 111 |
+
return rerank_with_flashrank(question, documents, top_k)
|
| 112 |
+
|
| 113 |
+
|
| 114 |
def evaluate_nli_groundedness(premise: str, hypothesis: str) -> Tuple[str, float]:
|
| 115 |
try:
|
| 116 |
model = get_nli_model()
|
app/engine/retriever.py
CHANGED
|
@@ -13,16 +13,49 @@ def get_retriever():
|
|
| 13 |
|
| 14 |
|
| 15 |
async def retrieve_documents(question: str, chat_history: list[dict] = None):
|
|
|
|
| 16 |
retriever = get_retriever()
|
| 17 |
documents = []
|
| 18 |
seen = set()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 19 |
variants = await query_variants(question, chat_history)
|
| 20 |
for query in variants:
|
| 21 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 22 |
key = document.metadata.get("chunk_id") or document.page_content[:120]
|
| 23 |
if key in seen:
|
| 24 |
continue
|
| 25 |
seen.add(key)
|
|
|
|
| 26 |
documents.append(document)
|
| 27 |
logger.info(
|
| 28 |
"retrieval completed query_count=%s returned_chunks=%s reranker_enabled=%s",
|
|
|
|
| 13 |
|
| 14 |
|
| 15 |
async def retrieve_documents(question: str, chat_history: list[dict] = None):
|
| 16 |
+
qdrant = open_vector_store()
|
| 17 |
retriever = get_retriever()
|
| 18 |
documents = []
|
| 19 |
seen = set()
|
| 20 |
+
|
| 21 |
+
try:
|
| 22 |
+
if hasattr(qdrant, "asimilarity_search_with_score"):
|
| 23 |
+
direct_results = await qdrant.asimilarity_search_with_score(question, k=settings.RETRIEVAL_TOP_K)
|
| 24 |
+
else:
|
| 25 |
+
direct_results = [(doc, 0.85) for doc in await retriever.ainvoke(question)]
|
| 26 |
+
except Exception:
|
| 27 |
+
direct_results = [(doc, 0.85) for doc in await retriever.ainvoke(question)]
|
| 28 |
+
|
| 29 |
+
for document, score in direct_results:
|
| 30 |
+
key = document.metadata.get("chunk_id") or document.page_content[:120]
|
| 31 |
+
if key not in seen:
|
| 32 |
+
seen.add(key)
|
| 33 |
+
document.metadata["similarity_score"] = float(score)
|
| 34 |
+
documents.append(document)
|
| 35 |
+
|
| 36 |
+
top_sim = max([d.metadata.get("similarity_score", 0.0) for d in documents] + [0.0])
|
| 37 |
+
if top_sim >= 0.65 and documents:
|
| 38 |
+
logger.info(f"Direct retrieval hit high similarity ({top_sim:.4f} >= 0.65). Skipping LLM query expansion!")
|
| 39 |
+
return documents[: settings.RETRIEVAL_TOP_K]
|
| 40 |
+
|
| 41 |
variants = await query_variants(question, chat_history)
|
| 42 |
for query in variants:
|
| 43 |
+
if query == question:
|
| 44 |
+
continue
|
| 45 |
+
try:
|
| 46 |
+
if hasattr(qdrant, "asimilarity_search_with_score"):
|
| 47 |
+
results = await qdrant.asimilarity_search_with_score(query, k=settings.RETRIEVAL_TOP_K)
|
| 48 |
+
else:
|
| 49 |
+
results = [(doc, 0.85) for doc in await retriever.ainvoke(query)]
|
| 50 |
+
except Exception:
|
| 51 |
+
results = [(doc, 0.85) for doc in await retriever.ainvoke(query)]
|
| 52 |
+
|
| 53 |
+
for document, score in results:
|
| 54 |
key = document.metadata.get("chunk_id") or document.page_content[:120]
|
| 55 |
if key in seen:
|
| 56 |
continue
|
| 57 |
seen.add(key)
|
| 58 |
+
document.metadata["similarity_score"] = float(score)
|
| 59 |
documents.append(document)
|
| 60 |
logger.info(
|
| 61 |
"retrieval completed query_count=%s returned_chunks=%s reranker_enabled=%s",
|
app/engine/semantic_cache.py
CHANGED
|
@@ -46,9 +46,17 @@ async def get_cached_answer(query: str, similarity_threshold: float = 0.92) -> O
|
|
| 46 |
|
| 47 |
if best_match:
|
| 48 |
logger.info(f"Semantic cache HIT (similarity: {best_sim:.4f} >= {similarity_threshold}) for query: '{query}'")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 49 |
return {
|
| 50 |
"answer": best_match["answer"],
|
| 51 |
-
"sources":
|
| 52 |
"confidence": best_match.get("confidence", 0.99),
|
| 53 |
"cached": True,
|
| 54 |
"similarity": best_sim
|
|
@@ -68,13 +76,27 @@ async def set_cached_answer(query: str, answer: str, sources: List[Any], confide
|
|
| 68 |
formatted_sources = []
|
| 69 |
for s in sources:
|
| 70 |
if isinstance(s, dict):
|
| 71 |
-
|
| 72 |
-
|
| 73 |
-
|
| 74 |
-
|
| 75 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 76 |
else:
|
| 77 |
-
formatted_sources.append(str(s))
|
| 78 |
|
| 79 |
key = f"semantic_cache:{abs(hash(query))}"
|
| 80 |
payload = {
|
|
|
|
| 46 |
|
| 47 |
if best_match:
|
| 48 |
logger.info(f"Semantic cache HIT (similarity: {best_sim:.4f} >= {similarity_threshold}) for query: '{query}'")
|
| 49 |
+
formatted_sources = []
|
| 50 |
+
for s in best_match.get("sources", []):
|
| 51 |
+
src_dict = dict(s) if isinstance(s, dict) else {"source": str(s)}
|
| 52 |
+
if "snippet" not in src_dict:
|
| 53 |
+
src_dict["snippet"] = src_dict.get("page_content", src_dict.get("source", best_match["answer"]))[:250]
|
| 54 |
+
if "source" not in src_dict:
|
| 55 |
+
src_dict["source"] = src_dict.get("doc_id", "cached_doc")
|
| 56 |
+
formatted_sources.append(src_dict)
|
| 57 |
return {
|
| 58 |
"answer": best_match["answer"],
|
| 59 |
+
"sources": formatted_sources,
|
| 60 |
"confidence": best_match.get("confidence", 0.99),
|
| 61 |
"cached": True,
|
| 62 |
"similarity": best_sim
|
|
|
|
| 76 |
formatted_sources = []
|
| 77 |
for s in sources:
|
| 78 |
if isinstance(s, dict):
|
| 79 |
+
src_dict = dict(s)
|
| 80 |
+
if "snippet" not in src_dict:
|
| 81 |
+
src_dict["snippet"] = src_dict.get("page_content", src_dict.get("source", str(s)))[:250]
|
| 82 |
+
if "source" not in src_dict:
|
| 83 |
+
src_dict["source"] = src_dict.get("doc_id", "cached_doc")
|
| 84 |
+
formatted_sources.append(src_dict)
|
| 85 |
+
elif hasattr(s, "dict") or hasattr(s, "model_dump"):
|
| 86 |
+
src_dict = s.model_dump() if hasattr(s, "model_dump") else s.dict()
|
| 87 |
+
if "snippet" not in src_dict:
|
| 88 |
+
src_dict["snippet"] = getattr(s, "page_content", str(s))[:250]
|
| 89 |
+
if "source" not in src_dict:
|
| 90 |
+
src_dict["source"] = getattr(s, "metadata", {}).get("source", getattr(s, "metadata", {}).get("doc_id", "cached_doc"))
|
| 91 |
+
formatted_sources.append(src_dict)
|
| 92 |
+
elif hasattr(s, "page_content"):
|
| 93 |
+
formatted_sources.append({
|
| 94 |
+
"source": getattr(s, "metadata", {}).get("source", getattr(s, "metadata", {}).get("doc_id", "cached_doc")),
|
| 95 |
+
"doc_id": getattr(s, "metadata", {}).get("doc_id", "cached_doc"),
|
| 96 |
+
"snippet": s.page_content[:250]
|
| 97 |
+
})
|
| 98 |
else:
|
| 99 |
+
formatted_sources.append({"source": "cached_doc", "snippet": str(s)[:250]})
|
| 100 |
|
| 101 |
key = f"semantic_cache:{abs(hash(query))}"
|
| 102 |
payload = {
|
app/graph/workflow.py
CHANGED
|
@@ -4,6 +4,7 @@ from langchain_core.prompts import PromptTemplate
|
|
| 4 |
from langchain_core.documents import Document
|
| 5 |
from langchain_openai import ChatOpenAI
|
| 6 |
from langgraph.graph import START, END, StateGraph
|
|
|
|
| 7 |
|
| 8 |
from app.core.config import settings
|
| 9 |
from app.core.logging import logger
|
|
@@ -11,6 +12,7 @@ from app.engine.context_builder import build_context, source_citations
|
|
| 11 |
from app.engine.retriever import retrieve_documents
|
| 12 |
from app.engine.reranker import rerank_documents, evaluate_nli_groundedness
|
| 13 |
|
|
|
|
| 14 |
class GraphState(TypedDict):
|
| 15 |
question: str
|
| 16 |
chat_history: List[dict]
|
|
@@ -20,8 +22,10 @@ class GraphState(TypedDict):
|
|
| 20 |
run_count: int
|
| 21 |
confidence_score: float
|
| 22 |
grounded: str
|
|
|
|
|
|
|
|
|
|
| 23 |
|
| 24 |
-
import httpx
|
| 25 |
|
| 26 |
_http_client = httpx.AsyncClient(
|
| 27 |
http2=True,
|
|
@@ -37,14 +41,7 @@ llm = ChatOpenAI(
|
|
| 37 |
default_headers={"HTTP-Referer": "https://localhost:3000", "X-Title": "Support Docs Copilot"},
|
| 38 |
http_async_client=_http_client,
|
| 39 |
)
|
| 40 |
-
|
| 41 |
-
model=getattr(settings, "FAST_LLM_MODEL", settings.LLM_MODEL),
|
| 42 |
-
temperature=0,
|
| 43 |
-
openai_api_key=settings.OPENROUTER_API_KEY,
|
| 44 |
-
openai_api_base=settings.OPENROUTER_BASE_URL,
|
| 45 |
-
default_headers={"HTTP-Referer": "https://localhost:3000", "X-Title": "Support Docs Copilot"},
|
| 46 |
-
http_async_client=_http_client,
|
| 47 |
-
)
|
| 48 |
llm_slow = ChatOpenAI(
|
| 49 |
model=getattr(settings, "SLOW_LLM_MODEL", settings.LLM_MODEL),
|
| 50 |
temperature=0,
|
|
@@ -54,75 +51,71 @@ llm_slow = ChatOpenAI(
|
|
| 54 |
http_async_client=_http_client,
|
| 55 |
)
|
| 56 |
|
|
|
|
| 57 |
async def retrieve(state: GraphState):
|
| 58 |
logger.info("NODE: RETRIEVE DOCS")
|
| 59 |
question = state["question"]
|
| 60 |
chat_history = state.get("chat_history", [])
|
| 61 |
run_count = state.get("run_count", 0)
|
| 62 |
documents = await retrieve_documents(question, chat_history)
|
| 63 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 64 |
|
| 65 |
async def grade_documents(state: GraphState):
|
| 66 |
-
logger.info("NODE: GRADE DOCUMENT RELEVANCE")
|
| 67 |
question = state["question"]
|
| 68 |
documents = state.get("documents", [])
|
| 69 |
|
| 70 |
-
reranked_docs = rerank_documents(question, documents, top_k=
|
| 71 |
if not reranked_docs:
|
| 72 |
return {"documents": []}
|
| 73 |
|
| 74 |
-
top_score = reranked_docs[0].metadata.get("rerank_score", -10.0)
|
| 75 |
-
if top_score >= 0.0:
|
| 76 |
-
logger.info(f"High confidence Cross-Encoder score ({top_score:.4f} >= 0.0). Skipping LLM grader.")
|
| 77 |
-
return {"documents": reranked_docs}
|
| 78 |
-
|
| 79 |
-
docs_text = "\n\n".join([f"[{idx+1}] ID: {d.metadata.get('doc_id', idx+1)}\nContent: {d.page_content}" for idx, d in enumerate(reranked_docs)])
|
| 80 |
-
prompt = PromptTemplate(
|
| 81 |
-
template="""You are a strict grader assessing relevance of retrieved documents to a user question.
|
| 82 |
-
User Question: {question}
|
| 83 |
-
|
| 84 |
-
Retrieved Documents:
|
| 85 |
-
{docs_text}
|
| 86 |
-
|
| 87 |
-
For each document [1] to [{count}], assess if it contains keywords or semantic meaning relevant to the question.
|
| 88 |
-
Return ONLY a JSON object with a key 'results' containing a list of objects: [{{"id": 1, "relevant": true}}, ...].""",
|
| 89 |
-
input_variables=["question", "docs_text", "count"],
|
| 90 |
-
)
|
| 91 |
-
grader = prompt | llm_json
|
| 92 |
-
result = await grader.ainvoke({"question": question, "docs_text": docs_text, "count": len(reranked_docs)})
|
| 93 |
-
|
| 94 |
filtered_docs = []
|
| 95 |
-
|
| 96 |
-
|
| 97 |
-
|
| 98 |
-
|
| 99 |
-
|
| 100 |
-
if r.get("relevant") is True or str(r.get("relevant")).lower() == "true":
|
| 101 |
-
idx_val = r.get("id")
|
| 102 |
-
if isinstance(idx_val, int) and 1 <= idx_val <= len(reranked_docs):
|
| 103 |
-
relevant_indices.add(idx_val - 1)
|
| 104 |
-
for i, d in enumerate(reranked_docs):
|
| 105 |
-
if i in relevant_indices:
|
| 106 |
-
filtered_docs.append(d)
|
| 107 |
-
except Exception as e:
|
| 108 |
-
logger.warning(f"Batch grading parse failed ({e}), keeping all {len(reranked_docs)} reranked docs.")
|
| 109 |
-
filtered_docs = reranked_docs
|
| 110 |
-
|
| 111 |
if not filtered_docs and reranked_docs:
|
| 112 |
-
top_score = reranked_docs[0].metadata.get("rerank_score", 0)
|
| 113 |
if top_score > 0.0:
|
| 114 |
filtered_docs = [reranked_docs[0]]
|
| 115 |
|
|
|
|
| 116 |
return {"documents": filtered_docs}
|
| 117 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 118 |
async def generate(state: GraphState):
|
| 119 |
logger.info("NODE: GENERATE ANSWER")
|
| 120 |
question = state["question"]
|
| 121 |
documents = state["documents"]
|
| 122 |
chat_history = state.get("chat_history", [])
|
|
|
|
| 123 |
run_count = state.get("run_count", 0) + 1
|
| 124 |
|
| 125 |
-
|
|
|
|
|
|
|
|
|
|
| 126 |
context = build_context(documents)
|
| 127 |
prompt = PromptTemplate(
|
| 128 |
template="""You are a Support Docs Copilot. Use only the retrieved context to answer the question concisely.
|
|
@@ -145,12 +138,6 @@ async def generate(state: GraphState):
|
|
| 145 |
generation = await rag_chain.ainvoke({"context": context, "question": question, "chat_history": history_str})
|
| 146 |
return {"generation": generation.content, "sources": source_citations(documents), "run_count": run_count}
|
| 147 |
|
| 148 |
-
async def decide_to_generate(state: GraphState):
|
| 149 |
-
if not state["documents"]:
|
| 150 |
-
logger.info("ROUTE: ALL DOCS IRRELEVANT")
|
| 151 |
-
return "end"
|
| 152 |
-
logger.info("ROUTE: RELEVANT DOCS FOUND")
|
| 153 |
-
return "generate"
|
| 154 |
|
| 155 |
async def evaluate_answer(state: GraphState):
|
| 156 |
logger.info("NODE: EVALUATE ANSWER")
|
|
@@ -162,6 +149,7 @@ async def evaluate_answer(state: GraphState):
|
|
| 162 |
|
| 163 |
return {"grounded": grade, "confidence_score": confidence}
|
| 164 |
|
|
|
|
| 165 |
async def check_hallucinations(state: GraphState):
|
| 166 |
run_count = state["run_count"]
|
| 167 |
|
|
@@ -177,6 +165,14 @@ async def check_hallucinations(state: GraphState):
|
|
| 177 |
logger.info("ROUTE: HALLUCINATION DETECTED")
|
| 178 |
return "regenerate"
|
| 179 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 180 |
def compile_workflow():
|
| 181 |
workflow = StateGraph(GraphState)
|
| 182 |
workflow.add_node("retrieve", retrieve)
|
|
@@ -184,7 +180,7 @@ def compile_workflow():
|
|
| 184 |
workflow.add_node("generate", generate)
|
| 185 |
workflow.add_node("evaluate_answer", evaluate_answer)
|
| 186 |
workflow.add_edge(START, "retrieve")
|
| 187 |
-
workflow.
|
| 188 |
workflow.add_conditional_edges("grade_documents", decide_to_generate, {"generate": "generate", "end": END})
|
| 189 |
workflow.add_edge("generate", "evaluate_answer")
|
| 190 |
workflow.add_conditional_edges("evaluate_answer", check_hallucinations, {"end": END, "regenerate": "generate"})
|
|
|
|
| 4 |
from langchain_core.documents import Document
|
| 5 |
from langchain_openai import ChatOpenAI
|
| 6 |
from langgraph.graph import START, END, StateGraph
|
| 7 |
+
import httpx
|
| 8 |
|
| 9 |
from app.core.config import settings
|
| 10 |
from app.core.logging import logger
|
|
|
|
| 12 |
from app.engine.retriever import retrieve_documents
|
| 13 |
from app.engine.reranker import rerank_documents, evaluate_nli_groundedness
|
| 14 |
|
| 15 |
+
|
| 16 |
class GraphState(TypedDict):
|
| 17 |
question: str
|
| 18 |
chat_history: List[dict]
|
|
|
|
| 22 |
run_count: int
|
| 23 |
confidence_score: float
|
| 24 |
grounded: str
|
| 25 |
+
summary: Optional[str]
|
| 26 |
+
optimistic_route: Optional[bool]
|
| 27 |
+
max_similarity: Optional[float]
|
| 28 |
|
|
|
|
| 29 |
|
| 30 |
_http_client = httpx.AsyncClient(
|
| 31 |
http2=True,
|
|
|
|
| 41 |
default_headers={"HTTP-Referer": "https://localhost:3000", "X-Title": "Support Docs Copilot"},
|
| 42 |
http_async_client=_http_client,
|
| 43 |
)
|
| 44 |
+
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 45 |
llm_slow = ChatOpenAI(
|
| 46 |
model=getattr(settings, "SLOW_LLM_MODEL", settings.LLM_MODEL),
|
| 47 |
temperature=0,
|
|
|
|
| 51 |
http_async_client=_http_client,
|
| 52 |
)
|
| 53 |
|
| 54 |
+
|
| 55 |
async def retrieve(state: GraphState):
|
| 56 |
logger.info("NODE: RETRIEVE DOCS")
|
| 57 |
question = state["question"]
|
| 58 |
chat_history = state.get("chat_history", [])
|
| 59 |
run_count = state.get("run_count", 0)
|
| 60 |
documents = await retrieve_documents(question, chat_history)
|
| 61 |
+
max_sim = max([d.metadata.get("similarity_score", 0.0) for d in documents] + [0.0])
|
| 62 |
+
optimistic = max_sim >= 0.82
|
| 63 |
+
if optimistic:
|
| 64 |
+
logger.info(f"OPTIMISTIC ROUTE TRIGGERED: Top similarity score {max_sim:.4f} >= 0.82")
|
| 65 |
+
return {
|
| 66 |
+
"documents": documents,
|
| 67 |
+
"sources": source_citations(documents),
|
| 68 |
+
"question": question,
|
| 69 |
+
"run_count": run_count,
|
| 70 |
+
"optimistic_route": optimistic,
|
| 71 |
+
"max_similarity": max_sim,
|
| 72 |
+
}
|
| 73 |
+
|
| 74 |
|
| 75 |
async def grade_documents(state: GraphState):
|
| 76 |
+
logger.info("NODE: GRADE DOCUMENT RELEVANCE (VIA HYBRID RERANKER)")
|
| 77 |
question = state["question"]
|
| 78 |
documents = state.get("documents", [])
|
| 79 |
|
| 80 |
+
reranked_docs = rerank_documents(question, documents, top_k=settings.RERANKER_TOP_N)
|
| 81 |
if not reranked_docs:
|
| 82 |
return {"documents": []}
|
| 83 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 84 |
filtered_docs = []
|
| 85 |
+
for doc in reranked_docs:
|
| 86 |
+
score = doc.metadata.get("relevance_score", doc.metadata.get("rerank_score", 1.0))
|
| 87 |
+
if score >= settings.MIN_RELEVANCE_SCORE:
|
| 88 |
+
filtered_docs.append(doc)
|
| 89 |
+
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 90 |
if not filtered_docs and reranked_docs:
|
| 91 |
+
top_score = reranked_docs[0].metadata.get("rerank_score", 0.0)
|
| 92 |
if top_score > 0.0:
|
| 93 |
filtered_docs = [reranked_docs[0]]
|
| 94 |
|
| 95 |
+
logger.info(f"Relevance grader filtered {len(reranked_docs)} docs down to {len(filtered_docs)} relevant docs (threshold >= {settings.MIN_RELEVANCE_SCORE}).")
|
| 96 |
return {"documents": filtered_docs}
|
| 97 |
|
| 98 |
+
|
| 99 |
+
async def decide_to_generate(state: GraphState):
|
| 100 |
+
if not state.get("documents"):
|
| 101 |
+
logger.info("ROUTE: ALL DOCS IRRELEVANT")
|
| 102 |
+
return "end"
|
| 103 |
+
logger.info("ROUTE: RELEVANT DOCS FOUND")
|
| 104 |
+
return "generate"
|
| 105 |
+
|
| 106 |
+
|
| 107 |
async def generate(state: GraphState):
|
| 108 |
logger.info("NODE: GENERATE ANSWER")
|
| 109 |
question = state["question"]
|
| 110 |
documents = state["documents"]
|
| 111 |
chat_history = state.get("chat_history", [])
|
| 112 |
+
summary = state.get("summary", "")
|
| 113 |
run_count = state.get("run_count", 0) + 1
|
| 114 |
|
| 115 |
+
history_lines = [f"{msg['role']}: {msg['content']}" for msg in chat_history[-6:]]
|
| 116 |
+
if summary:
|
| 117 |
+
history_lines.insert(0, summary if summary.startswith("System Summary:") else f"System Summary: {summary}")
|
| 118 |
+
history_str = "\n".join(history_lines)
|
| 119 |
context = build_context(documents)
|
| 120 |
prompt = PromptTemplate(
|
| 121 |
template="""You are a Support Docs Copilot. Use only the retrieved context to answer the question concisely.
|
|
|
|
| 138 |
generation = await rag_chain.ainvoke({"context": context, "question": question, "chat_history": history_str})
|
| 139 |
return {"generation": generation.content, "sources": source_citations(documents), "run_count": run_count}
|
| 140 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 141 |
|
| 142 |
async def evaluate_answer(state: GraphState):
|
| 143 |
logger.info("NODE: EVALUATE ANSWER")
|
|
|
|
| 149 |
|
| 150 |
return {"grounded": grade, "confidence_score": confidence}
|
| 151 |
|
| 152 |
+
|
| 153 |
async def check_hallucinations(state: GraphState):
|
| 154 |
run_count = state["run_count"]
|
| 155 |
|
|
|
|
| 165 |
logger.info("ROUTE: HALLUCINATION DETECTED")
|
| 166 |
return "regenerate"
|
| 167 |
|
| 168 |
+
|
| 169 |
+
async def decide_optimistic_or_grade(state: GraphState):
|
| 170 |
+
if state.get("optimistic_route") and state.get("documents"):
|
| 171 |
+
logger.info(f"OPTIMISTIC STREAMING: High similarity ({state.get('max_similarity', 0.0):.4f} >= 0.82). Skipping LLM grader node!")
|
| 172 |
+
return "generate"
|
| 173 |
+
return "grade_documents"
|
| 174 |
+
|
| 175 |
+
|
| 176 |
def compile_workflow():
|
| 177 |
workflow = StateGraph(GraphState)
|
| 178 |
workflow.add_node("retrieve", retrieve)
|
|
|
|
| 180 |
workflow.add_node("generate", generate)
|
| 181 |
workflow.add_node("evaluate_answer", evaluate_answer)
|
| 182 |
workflow.add_edge(START, "retrieve")
|
| 183 |
+
workflow.add_conditional_edges("retrieve", decide_optimistic_or_grade, {"generate": "generate", "grade_documents": "grade_documents"})
|
| 184 |
workflow.add_conditional_edges("grade_documents", decide_to_generate, {"generate": "generate", "end": END})
|
| 185 |
workflow.add_edge("generate", "evaluate_answer")
|
| 186 |
workflow.add_conditional_edges("evaluate_answer", check_hallucinations, {"end": END, "regenerate": "generate"})
|
app/guardrails/input.py
CHANGED
|
@@ -19,7 +19,6 @@ PROMPT_INJECTION_PATTERNS = [
|
|
| 19 |
"disregard instructions",
|
| 20 |
"reveal hidden",
|
| 21 |
"you are now an arbitrary",
|
| 22 |
-
"dan",
|
| 23 |
"do anything now",
|
| 24 |
"ignore all constraints"
|
| 25 |
]
|
|
|
|
| 19 |
"disregard instructions",
|
| 20 |
"reveal hidden",
|
| 21 |
"you are now an arbitrary",
|
|
|
|
| 22 |
"do anything now",
|
| 23 |
"ignore all constraints"
|
| 24 |
]
|
app/main.py
CHANGED
|
@@ -19,11 +19,14 @@ from app.core.config import settings
|
|
| 19 |
from app.core.dependencies import check_openrouter, check_qdrant
|
| 20 |
from app.core.errors import CopilotError
|
| 21 |
from app.core.logging import configure_logging, logger, request_id_var
|
|
|
|
| 22 |
from app.engine.document_registry import load_registry
|
| 23 |
from app.engine.ingestion import delete_indexed_document, ingest_documents, reset_index
|
| 24 |
from app.engine.context_builder import build_context, format_sources
|
| 25 |
-
|
|
|
|
| 26 |
from app.engine.query_transform import condense_query
|
|
|
|
| 27 |
from app.engine.semantic_cache import get_cached_answer, set_cached_answer
|
| 28 |
from app.graph.workflow import compile_workflow
|
| 29 |
from app.guardrails.input import async_enforce_rate_limit, enforce_rate_limit, validate_query
|
|
@@ -42,6 +45,47 @@ if settings.LANGCHAIN_TRACING_V2 and settings.LANGCHAIN_API_KEY:
|
|
| 42 |
|
| 43 |
app = FastAPI(title=settings.PROJECT_NAME)
|
| 44 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 45 |
app.add_middleware(
|
| 46 |
CORSMiddleware,
|
| 47 |
allow_origins=["*"],
|
|
@@ -97,9 +141,13 @@ class IngestionRequest(BaseModel):
|
|
| 97 |
class FeedbackRequest(BaseModel):
|
| 98 |
query: str
|
| 99 |
answer: str
|
| 100 |
-
is_positive: bool
|
| 101 |
comments: str | None = None
|
| 102 |
|
|
|
|
|
|
|
|
|
|
|
|
|
| 103 |
@app.get("/health")
|
| 104 |
async def health_endpoint():
|
| 105 |
return {"status": "ok", "project": settings.PROJECT_NAME}
|
|
@@ -219,9 +267,10 @@ async def chat_endpoint(request: ChatRequest, http_request: Request, user: UserC
|
|
| 219 |
session_id = request.session_id or str(uuid.uuid4())
|
| 220 |
chat_history = request.chat_history
|
| 221 |
if not chat_history:
|
| 222 |
-
chat_history = await get_session_history(user.user_id, session_id)
|
|
|
|
| 223 |
|
| 224 |
-
standalone_query = await
|
| 225 |
|
| 226 |
cached = await get_cached_answer(standalone_query)
|
| 227 |
if cached:
|
|
@@ -230,7 +279,7 @@ async def chat_endpoint(request: ChatRequest, http_request: Request, user: UserC
|
|
| 230 |
await add_session_message(user.user_id, session_id, "assistant", cached["answer"], cached.get("sources", []), cached.get("confidence", 0.99))
|
| 231 |
return ChatResponse(query=request.query, answer=cached["answer"], sources=cached.get("sources", []), confidence=cached.get("confidence", 0.99), session_id=session_id)
|
| 232 |
|
| 233 |
-
initial_state = {"question": standalone_query, "chat_history": chat_history, "run_count": 0}
|
| 234 |
try:
|
| 235 |
with timed_stage(metrics, "rag_workflow"):
|
| 236 |
final_state = await rag_agent.ainvoke(initial_state)
|
|
@@ -266,9 +315,10 @@ async def chat_stream_endpoint(request: ChatRequest, http_request: Request, user
|
|
| 266 |
session_id = request.session_id or str(uuid.uuid4())
|
| 267 |
chat_history = request.chat_history
|
| 268 |
if not chat_history:
|
| 269 |
-
chat_history = await get_session_history(user.user_id, session_id)
|
|
|
|
| 270 |
|
| 271 |
-
standalone_query = await
|
| 272 |
|
| 273 |
async def token_generator():
|
| 274 |
try:
|
|
@@ -279,12 +329,9 @@ async def chat_stream_endpoint(request: ChatRequest, http_request: Request, user
|
|
| 279 |
await add_session_message(user.user_id, session_id, "user", request.query)
|
| 280 |
await add_session_message(user.user_id, session_id, "assistant", cached["answer"], cached.get("sources", []), cached.get("confidence", 0.99))
|
| 281 |
yield cached["answer"]
|
| 282 |
-
sources_list = cached.get("sources", [])
|
| 283 |
-
if sources_list:
|
| 284 |
-
yield f"\n\n{format_sources(sources_list)}" if isinstance(sources_list, list) and sources_list and hasattr(sources_list[0], 'metadata') else f"\n\n{sources_list}"
|
| 285 |
return
|
| 286 |
|
| 287 |
-
initial_state = {"question": standalone_query, "chat_history": chat_history, "run_count": 0}
|
| 288 |
documents = []
|
| 289 |
sources_text = ""
|
| 290 |
grounded_result = "yes"
|
|
@@ -292,11 +339,18 @@ async def chat_stream_endpoint(request: ChatRequest, http_request: Request, user
|
|
| 292 |
streamed_text = ""
|
| 293 |
|
| 294 |
with timed_stage(metrics, "rag_workflow_stream"):
|
|
|
|
| 295 |
async for event in rag_agent.astream_events(initial_state, version="v2"):
|
| 296 |
kind = event["event"]
|
| 297 |
node_name = event.get("metadata", {}).get("langgraph_node", "")
|
| 298 |
|
| 299 |
if kind == "on_chat_model_stream" and node_name == "generate":
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 300 |
chunk = event["data"]["chunk"]
|
| 301 |
if chunk and getattr(chunk, "content", None):
|
| 302 |
has_streamed_tokens = True
|
|
@@ -325,11 +379,6 @@ async def chat_stream_endpoint(request: ChatRequest, http_request: Request, user
|
|
| 325 |
yield "\n\n🚨 **[CANCELLED: This response violated safety guidelines and has been retracted.]**"
|
| 326 |
return
|
| 327 |
|
| 328 |
-
if sources_text:
|
| 329 |
-
yield f"\n\n{sources_text}"
|
| 330 |
-
elif documents:
|
| 331 |
-
yield format_sources(documents)
|
| 332 |
-
|
| 333 |
if has_streamed_tokens and documents and str(grounded_result).lower() != "no":
|
| 334 |
redacted_text = redact_sensitive_data(streamed_text)
|
| 335 |
await set_cached_answer(standalone_query, redacted_text, documents, 0.98)
|
|
@@ -356,8 +405,37 @@ async def get_session_messages_endpoint(session_id: str, user: UserContext = Dep
|
|
| 356 |
messages = await get_session_history(user.user_id, session_id, limit=50)
|
| 357 |
return {"session_id": session_id, "messages": messages}
|
| 358 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 359 |
@app.delete("/api/v1/sessions/{session_id}")
|
| 360 |
async def delete_session_endpoint(session_id: str, user: UserContext = Depends(resolve_user)):
|
| 361 |
success = await delete_session(user.user_id, session_id)
|
| 362 |
return {"status": "ok" if success else "error", "session_id": session_id}
|
| 363 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 19 |
from app.core.dependencies import check_openrouter, check_qdrant
|
| 20 |
from app.core.errors import CopilotError
|
| 21 |
from app.core.logging import configure_logging, logger, request_id_var
|
| 22 |
+
from app.core.queue import get_redis_client
|
| 23 |
from app.engine.document_registry import load_registry
|
| 24 |
from app.engine.ingestion import delete_indexed_document, ingest_documents, reset_index
|
| 25 |
from app.engine.context_builder import build_context, format_sources
|
| 26 |
+
import csv
|
| 27 |
+
from app.engine.memory import add_session_message, get_session_history, list_user_sessions, delete_session, get_session_summary, list_all_sessions
|
| 28 |
from app.engine.query_transform import condense_query
|
| 29 |
+
from app.engine.retriever import retrieve_documents
|
| 30 |
from app.engine.semantic_cache import get_cached_answer, set_cached_answer
|
| 31 |
from app.graph.workflow import compile_workflow
|
| 32 |
from app.guardrails.input import async_enforce_rate_limit, enforce_rate_limit, validate_query
|
|
|
|
| 45 |
|
| 46 |
app = FastAPI(title=settings.PROJECT_NAME)
|
| 47 |
|
| 48 |
+
@app.on_event("startup")
|
| 49 |
+
async def startup_faq_prewarming():
|
| 50 |
+
try:
|
| 51 |
+
csv_path = Path("datasets/golden_qa.csv")
|
| 52 |
+
if csv_path.exists():
|
| 53 |
+
logger.info("PRE-WARMING SEMANTIC CACHE: Seeding FAQ entries from golden_qa.csv...")
|
| 54 |
+
with open(csv_path, mode="r", encoding="utf-8") as f:
|
| 55 |
+
reader = csv.DictReader(f)
|
| 56 |
+
count = 0
|
| 57 |
+
for row in reader:
|
| 58 |
+
question = row.get("question", "").strip()
|
| 59 |
+
answer = row.get("expected_answer", "").strip()
|
| 60 |
+
sources_raw = row.get("expected_sources", "").strip()
|
| 61 |
+
if question and answer:
|
| 62 |
+
sources = [{"doc_id": sources_raw, "source": sources_raw, "snippet": answer}] if sources_raw else []
|
| 63 |
+
await set_cached_answer(question, answer, sources, confidence=0.99)
|
| 64 |
+
count += 1
|
| 65 |
+
logger.info(f"PRE-WARMING COMPLETE: Successfully seeded {count} FAQ entries into Redis vector cache.")
|
| 66 |
+
except Exception as e:
|
| 67 |
+
logger.warning(f"FAQ pre-warming failed or skipped: {e}")
|
| 68 |
+
|
| 69 |
+
async def resolve_query_speculative(query: str, chat_history: list, summary: str):
|
| 70 |
+
speculative_docs = []
|
| 71 |
+
if chat_history:
|
| 72 |
+
condense_task = asyncio.create_task(condense_query(query, chat_history, summary=summary))
|
| 73 |
+
retrieval_task = asyncio.create_task(retrieve_documents(query, chat_history))
|
| 74 |
+
results = await asyncio.gather(condense_task, retrieval_task, return_exceptions=True)
|
| 75 |
+
|
| 76 |
+
standalone_query = query if isinstance(results[0], Exception) else results[0]
|
| 77 |
+
raw_docs = [] if isinstance(results[1], Exception) else results[1]
|
| 78 |
+
|
| 79 |
+
if raw_docs:
|
| 80 |
+
top_sim = max([d.metadata.get("similarity_score", 0.0) for d in raw_docs] + [0.0])
|
| 81 |
+
if top_sim >= 0.85:
|
| 82 |
+
logger.info(f"SPECULATIVE RETRIEVAL HIT: Raw query '{query}' matched with top similarity {top_sim:.4f} >= 0.85!")
|
| 83 |
+
speculative_docs = raw_docs
|
| 84 |
+
else:
|
| 85 |
+
standalone_query = await condense_query(query, chat_history, summary=summary)
|
| 86 |
+
|
| 87 |
+
return standalone_query, speculative_docs
|
| 88 |
+
|
| 89 |
app.add_middleware(
|
| 90 |
CORSMiddleware,
|
| 91 |
allow_origins=["*"],
|
|
|
|
| 141 |
class FeedbackRequest(BaseModel):
|
| 142 |
query: str
|
| 143 |
answer: str
|
| 144 |
+
is_positive: bool = True
|
| 145 |
comments: str | None = None
|
| 146 |
|
| 147 |
+
class InterveneRequest(BaseModel):
|
| 148 |
+
message: str
|
| 149 |
+
role: str = "assistant"
|
| 150 |
+
|
| 151 |
@app.get("/health")
|
| 152 |
async def health_endpoint():
|
| 153 |
return {"status": "ok", "project": settings.PROJECT_NAME}
|
|
|
|
| 267 |
session_id = request.session_id or str(uuid.uuid4())
|
| 268 |
chat_history = request.chat_history
|
| 269 |
if not chat_history:
|
| 270 |
+
chat_history = await get_session_history(user.user_id, session_id, limit=6)
|
| 271 |
+
summary = await get_session_summary(user.user_id, session_id)
|
| 272 |
|
| 273 |
+
standalone_query, speculative_docs = await resolve_query_speculative(request.query, chat_history, summary=summary)
|
| 274 |
|
| 275 |
cached = await get_cached_answer(standalone_query)
|
| 276 |
if cached:
|
|
|
|
| 279 |
await add_session_message(user.user_id, session_id, "assistant", cached["answer"], cached.get("sources", []), cached.get("confidence", 0.99))
|
| 280 |
return ChatResponse(query=request.query, answer=cached["answer"], sources=cached.get("sources", []), confidence=cached.get("confidence", 0.99), session_id=session_id)
|
| 281 |
|
| 282 |
+
initial_state = {"question": standalone_query, "chat_history": chat_history, "summary": summary, "run_count": 0, "documents": speculative_docs}
|
| 283 |
try:
|
| 284 |
with timed_stage(metrics, "rag_workflow"):
|
| 285 |
final_state = await rag_agent.ainvoke(initial_state)
|
|
|
|
| 315 |
session_id = request.session_id or str(uuid.uuid4())
|
| 316 |
chat_history = request.chat_history
|
| 317 |
if not chat_history:
|
| 318 |
+
chat_history = await get_session_history(user.user_id, session_id, limit=6)
|
| 319 |
+
summary = await get_session_summary(user.user_id, session_id)
|
| 320 |
|
| 321 |
+
standalone_query, speculative_docs = await resolve_query_speculative(request.query, chat_history, summary=summary)
|
| 322 |
|
| 323 |
async def token_generator():
|
| 324 |
try:
|
|
|
|
| 329 |
await add_session_message(user.user_id, session_id, "user", request.query)
|
| 330 |
await add_session_message(user.user_id, session_id, "assistant", cached["answer"], cached.get("sources", []), cached.get("confidence", 0.99))
|
| 331 |
yield cached["answer"]
|
|
|
|
|
|
|
|
|
|
| 332 |
return
|
| 333 |
|
| 334 |
+
initial_state = {"question": standalone_query, "chat_history": chat_history, "summary": summary, "run_count": 0, "documents": speculative_docs}
|
| 335 |
documents = []
|
| 336 |
sources_text = ""
|
| 337 |
grounded_result = "yes"
|
|
|
|
| 339 |
streamed_text = ""
|
| 340 |
|
| 341 |
with timed_stage(metrics, "rag_workflow_stream"):
|
| 342 |
+
redis = await get_redis_client()
|
| 343 |
async for event in rag_agent.astream_events(initial_state, version="v2"):
|
| 344 |
kind = event["event"]
|
| 345 |
node_name = event.get("metadata", {}).get("langgraph_node", "")
|
| 346 |
|
| 347 |
if kind == "on_chat_model_stream" and node_name == "generate":
|
| 348 |
+
if await redis.exists(f"session:{session_id}:terminate"):
|
| 349 |
+
await redis.delete(f"session:{session_id}:terminate")
|
| 350 |
+
yield "\n\n🛑 **[TERMINATED BY USER: Generation was stopped.]**"
|
| 351 |
+
await add_session_message(user.user_id, session_id, "user", request.query)
|
| 352 |
+
await add_session_message(user.user_id, session_id, "assistant", streamed_text + "\n\n🛑 [TERMINATED BY USER]", documents or [], 0.0)
|
| 353 |
+
return
|
| 354 |
chunk = event["data"]["chunk"]
|
| 355 |
if chunk and getattr(chunk, "content", None):
|
| 356 |
has_streamed_tokens = True
|
|
|
|
| 379 |
yield "\n\n🚨 **[CANCELLED: This response violated safety guidelines and has been retracted.]**"
|
| 380 |
return
|
| 381 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 382 |
if has_streamed_tokens and documents and str(grounded_result).lower() != "no":
|
| 383 |
redacted_text = redact_sensitive_data(streamed_text)
|
| 384 |
await set_cached_answer(standalone_query, redacted_text, documents, 0.98)
|
|
|
|
| 405 |
messages = await get_session_history(user.user_id, session_id, limit=50)
|
| 406 |
return {"session_id": session_id, "messages": messages}
|
| 407 |
|
| 408 |
+
@app.post("/api/v1/sessions/{session_id}/terminate")
|
| 409 |
+
async def terminate_session_endpoint(session_id: str, user: UserContext = Depends(resolve_user)):
|
| 410 |
+
redis = await get_redis_client()
|
| 411 |
+
await redis.setex(f"session:{session_id}:terminate", 60, "1")
|
| 412 |
+
return {"status": "terminated", "session_id": session_id}
|
| 413 |
+
|
| 414 |
@app.delete("/api/v1/sessions/{session_id}")
|
| 415 |
async def delete_session_endpoint(session_id: str, user: UserContext = Depends(resolve_user)):
|
| 416 |
success = await delete_session(user.user_id, session_id)
|
| 417 |
return {"status": "ok" if success else "error", "session_id": session_id}
|
| 418 |
|
| 419 |
+
@app.get("/api/v1/admin/sessions")
|
| 420 |
+
async def admin_list_sessions_endpoint(user: UserContext = Depends(resolve_user)):
|
| 421 |
+
if user.role != "admin":
|
| 422 |
+
raise CopilotError("Admin privileges required", status_code=403)
|
| 423 |
+
sessions = await list_all_sessions(limit=50)
|
| 424 |
+
return {"sessions": sessions}
|
| 425 |
+
|
| 426 |
+
@app.get("/api/v1/admin/sessions/{user_id}/{session_id}/messages")
|
| 427 |
+
async def admin_get_session_messages_endpoint(user_id: str, session_id: str, user: UserContext = Depends(resolve_user)):
|
| 428 |
+
if user.role != "admin":
|
| 429 |
+
raise CopilotError("Admin privileges required", status_code=403)
|
| 430 |
+
messages = await get_session_history(user_id, session_id, limit=50)
|
| 431 |
+
summary = await get_session_summary(user_id, session_id)
|
| 432 |
+
return {"session_id": session_id, "user_id": user_id, "messages": messages, "summary": summary}
|
| 433 |
+
|
| 434 |
+
@app.post("/api/v1/admin/sessions/{user_id}/{session_id}/message")
|
| 435 |
+
async def admin_intervene_message_endpoint(user_id: str, session_id: str, request: InterveneRequest, user: UserContext = Depends(resolve_user)):
|
| 436 |
+
if user.role != "admin":
|
| 437 |
+
raise CopilotError("Admin privileges required", status_code=403)
|
| 438 |
+
await add_session_message(user_id, session_id, request.role, request.message)
|
| 439 |
+
return {"status": "ok", "message": "Intervention message injected."}
|
| 440 |
+
|
| 441 |
+
|
app/tests/eval_rag.py
CHANGED
|
@@ -71,8 +71,8 @@ async def run_local_evaluation(
|
|
| 71 |
fast_llm = ChatOpenAI(
|
| 72 |
model=getattr(settings, "FAST_LLM_MODEL", settings.LLM_MODEL),
|
| 73 |
temperature=0,
|
| 74 |
-
openai_api_key=settings.OPENROUTER_API_KEY,
|
| 75 |
-
openai_api_base=settings.OPENROUTER_BASE_URL,
|
| 76 |
default_headers={"HTTP-Referer": "https://localhost:3000", "X-Title": "Support Docs Copilot"},
|
| 77 |
)
|
| 78 |
slow_llm = ChatOpenAI(
|
|
|
|
| 71 |
fast_llm = ChatOpenAI(
|
| 72 |
model=getattr(settings, "FAST_LLM_MODEL", settings.LLM_MODEL),
|
| 73 |
temperature=0,
|
| 74 |
+
openai_api_key=getattr(settings, "FAST_LLM_API_KEY", "") or settings.OPENROUTER_API_KEY,
|
| 75 |
+
openai_api_base=getattr(settings, "FAST_LLM_BASE_URL", "") or settings.OPENROUTER_BASE_URL,
|
| 76 |
default_headers={"HTTP-Referer": "https://localhost:3000", "X-Title": "Support Docs Copilot"},
|
| 77 |
)
|
| 78 |
slow_llm = ChatOpenAI(
|
requirements.txt
CHANGED
|
@@ -9,7 +9,7 @@ langchain-text-splitters==0.2.4
|
|
| 9 |
langchain-openai==0.1.22
|
| 10 |
langchain-qdrant==0.1.4
|
| 11 |
qdrant-client==1.10.1
|
| 12 |
-
fastembed=
|
| 13 |
langgraph==0.2.76
|
| 14 |
guardrails-ai==0.5.0
|
| 15 |
ragas==0.1.21
|
|
@@ -34,5 +34,7 @@ arq>=0.26.0
|
|
| 34 |
redis>=5.0.0
|
| 35 |
sentence-transformers>=2.7.0
|
| 36 |
torch==2.2.2+cpu
|
| 37 |
-
transformers
|
|
|
|
|
|
|
| 38 |
|
|
|
|
| 9 |
langchain-openai==0.1.22
|
| 10 |
langchain-qdrant==0.1.4
|
| 11 |
qdrant-client==1.10.1
|
| 12 |
+
fastembed>=0.4.1
|
| 13 |
langgraph==0.2.76
|
| 14 |
guardrails-ai==0.5.0
|
| 15 |
ragas==0.1.21
|
|
|
|
| 34 |
redis>=5.0.0
|
| 35 |
sentence-transformers>=2.7.0
|
| 36 |
torch==2.2.2+cpu
|
| 37 |
+
transformers==4.39.3
|
| 38 |
+
cohere>=7.0.0
|
| 39 |
+
flashrank>=0.2.8
|
| 40 |
|