Cognitive-rag / backend / tests / test_chat_answer_cleanup.py
test_chat_answer_cleanup.py
Raw
from __future__ import annotations

import asyncio
from types import SimpleNamespace

import backend.config as backend_config
import backend.routes_chat as routes_chat


class _FakeEmbedClient:
    def embed_query(self, _text: str) -> list[float]:
        return [0.01, 0.02, 0.03]


class _FakeEmbeddingQueryModule:
    @staticmethod
    def load_client(*_args, **_kwargs) -> _FakeEmbedClient:
        return _FakeEmbedClient()


class _FakeRetrieverModule:
    @staticmethod
    def get_unified_context_for_llm(**_kwargs) -> dict:
        return {
            "document_chunks": [
                {
                    "chunk_id": "doc_chunk_1",
                    "doc_id": "project-proposal-streamcore-pdf__027b5525",
                    "document_name": "Project Proposal StreamCore.pdf",
                    "text": "Deliver StreamLoader that reads long files by memory mapping.",
                    "section": "page_1",
                    "page": 1,
                    "score": 2.0,
                    "source_path": "Project Proposal StreamCore.pdf",
                }
            ],
            "related_chunks": [],
            "memory_chunks": [
                {
                    "chunk_id": "mem_chunk_1",
                    "doc_id": "shared_memory",
                    "text": 'Query: "old"\nAnswer: old memory text that should not leak.',
                    "score": 1.2,
                    "section": "conversation_memory",
                    "page": 1,
                }
            ],
            "stats": {"doc_count": 1, "memory_count": 1, "related_count": 0},
            "doc_count": 1,
            "memory_count": 1,
            "related_count": 0,
            "memory_raw_count": 1,
            "memory_dropped_count": 0,
        }

    @staticmethod
    def retrieve_dense(**_kwargs) -> dict:
        return {
            "retrieved_chunks": [
                {
                    "chunk_id": "doc_chunk_1",
                    "doc_id": "project-proposal-streamcore-pdf__027b5525",
                    "text": "Deliver StreamLoader that reads long files by memory mapping.",
                    "score": 2.0,
                }
            ]
        }

    @staticmethod
    def retrieve_sparse(**_kwargs) -> dict:
        return {
            "retrieved_chunks": [
                {
                    "chunk_id": "doc_chunk_1",
                    "doc_id": "project-proposal-streamcore-pdf__027b5525",
                    "text": "Deliver StreamLoader that reads long files by memory mapping.",
                    "score": 1.8,
                }
            ]
        }

    @staticmethod
    def rerank(_query: str, candidates: list[dict], top_n: int = 5) -> list[dict]:
        return candidates[:top_n]


class _FakeGenerationModule:
    @staticmethod
    def generate_answer(input_data: dict) -> str:
        query = input_data["query"]
        return f"""
Query:
    "{query}"

Answer:
Answer:
1. Deliver StreamLoader that reads long files by memory mapping.
2. Deliver StreamTrainer for chunk-based training.
2. Deliver StreamTrainer for chunk-based training.
[shared_memory] Query: "old"
[shared_memory] Answer: old memory answer fragment

Evidence:
    Deliver StreamLoader that reads long files by memory mapping.
""".strip()


class _FakeMemoryModule:
    @staticmethod
    def parse_LLM_output(_raw_output: str) -> dict:
        return {
            "answer": 'Answer:\n[shared_memory] Query: "old"\nrepeated parser answer',
            "evidence": [
                "Deliver StreamLoader that reads long files by memory mapping.",
                'Query: "old"\nAnswer: old memory text that should not leak.',
            ],
        }

    @staticmethod
    def store_memory_sync(*_args, **_kwargs):
        return None


class _FakeIndex:
    @staticmethod
    def query(**_kwargs) -> dict:
        return {"matches": []}


class _FakePineconeClient:
    @staticmethod
    def get_index() -> _FakeIndex:
        return _FakeIndex()

    @staticmethod
    def retrieve_similar(*_args, **_kwargs) -> dict:
        return {"retrieved_chunks": []}


class _FakePineconeModule:
    @staticmethod
    def get_client() -> _FakePineconeClient:
        return _FakePineconeClient()


class _FakeVectordbConfig:
    RERANKER_ENABLED = False
    RERANKER_TOP_N = 10
    CONVERSATION_MEMORY_NAMESPACE = "conversation_memory"


def test_chat_answer_is_clean_and_sources_preserved(monkeypatch):
    modules = {
        "embedding_query": _FakeEmbeddingQueryModule,
        "vectordb_retriever": _FakeRetrieverModule,
        "generation": _FakeGenerationModule,
        "memory": _FakeMemoryModule,
        "vectordb_pinecone": _FakePineconeModule,
        "vectordb_config": _FakeVectordbConfig,
    }

    monkeypatch.setattr(backend_config, "get_stage_modules", lambda: modules)
    monkeypatch.setattr(routes_chat, "resolve_user", lambda _ip: "user_1")
    monkeypatch.setattr(routes_chat, "audit_event", lambda *_args, **_kwargs: None)
    monkeypatch.setattr(
        routes_chat,
        "write_prompt_reports",
        lambda **_kwargs: {"retrieval_report": "x", "generation_report": "y"},
    )

    env_map = {
        "TOP_K": 5,
        "MEMORY_TOP_K": 3,
        "RERANKER_TOP_N": 5,
    }
    monkeypatch.setattr(
        backend_config,
        "env_int",
        lambda name, default=0: int(env_map.get(name, default)),
    )
    monkeypatch.setattr(
        backend_config,
        "env_value",
        lambda _name, default="": default,
    )

    request = SimpleNamespace(client=SimpleNamespace(host="127.0.0.1"))
    payload = routes_chat.ChatRequest(
        query="tell me what are the objectives list them all numbered and the timeline for the twenty weeks",
        debug=False,
        use_multihop=False,
    )
    response = asyncio.run(routes_chat.chat(request, payload))

    assert "sources" in response and response["sources"]
    assert response["sources"][0]["text"] == "Deliver StreamLoader that reads long files by memory mapping."
    assert all(item.get("doc_id") != "shared_memory" for item in response["sources"])
    assert all(item.get("section") != "conversation_memory" for item in response["sources"])

    answer = response["answer"]
    assert "Answer:" not in answer
    assert "Query:" not in answer
    assert answer.count("Evidence used:") == 1
    assert "1. Deliver StreamLoader that reads long files by memory mapping. [Project Proposal StreamCore.pdf" in answer
    assert "[shared_memory]" not in answer
    assert "conversation_memory" not in answer
    assert "old memory answer fragment" not in answer
    assert answer.count("2. Deliver StreamTrainer for chunk-based training.") == 1