diff --git a/VERSION b/VERSION index 0383441..7f20734 100644 --- a/VERSION +++ b/VERSION @@ -1 +1 @@ -0.9.5 \ No newline at end of file +1.0.1 \ No newline at end of file diff --git a/docker-compose-network-github-image.yml b/docker-compose-network-github-image.yml index ee25641..d206bb3 100644 --- a/docker-compose-network-github-image.yml +++ b/docker-compose-network-github-image.yml @@ -1,6 +1,6 @@ services: concept-graphs-api: - image: ghcr.io/onto-med/concept-graphs/concept-graphs-api:1.0.0 + image: ghcr.io/onto-med/concept-graphs/concept-graphs-api:1.0.1 restart: unless-stopped ports: - 9007:9007 diff --git a/docker-compose-network.yml b/docker-compose-network.yml index ef8c5b9..2b546a9 100644 --- a/docker-compose-network.yml +++ b/docker-compose-network.yml @@ -1,6 +1,6 @@ services: concept-graphs-api: - image: imise/top/concept-graphs-api:1.0.0 + image: imise/top/concept-graphs-api:1.0.1 restart: unless-stopped ports: - 9007:9007 diff --git a/docker-compose.yml b/docker-compose.yml index 27f29e3..caef56d 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -1,6 +1,6 @@ services: concept-graphs-api: - image: imise/top/concept-graphs-api:1.0.0 + image: imise/top/concept-graphs-api:1.0.1 restart: unless-stopped ports: - 9007:9007 diff --git a/main.py b/main.py index 7189cd9..6d82a95 100644 --- a/main.py +++ b/main.py @@ -1,4 +1,5 @@ import logging +import os import pathlib import flask @@ -15,7 +16,15 @@ def configure_logging(logging_setup_tuples: list[tuple] | None = None) -> None: - """Configure application logging defaults.""" + """Configure application logging defaults. + + ``LOG_LEVEL`` controls application/root logging and defaults to ``INFO`` so + operational messages such as RAG indexing progress are visible in container + logs. Noisy dependency loggers are still kept at warning level by default. + """ + log_level_name = os.getenv("LOG_LEVEL", "INFO").upper() + log_level = getattr(logging, log_level_name, logging.INFO) + if logging_setup_tuples is None: logging_setup_tuples = [ ("werkzeug", logging.WARN), @@ -26,9 +35,11 @@ def configure_logging(logging_setup_tuples: list[tuple] | None = None) -> None: logging.getLogger(logger_name).setLevel(level) root_logger = logging.getLogger() + root_logger.setLevel(log_level) root_logger.propagate = False if root_logger.hasHandlers(): root_logger.handlers.clear() + flask.logging.default_handler.setLevel(log_level) root_logger.addHandler(flask.logging.default_handler) diff --git a/pyproject.toml b/pyproject.toml index 0b4d9aa..3819073 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "concept-graphs" -version = "1.0.0" +version = "1.0.1" description = "" authors = [ { name = "Franz Matthies", email = "franz.matthies@imise.uni-leipzig.de" }, @@ -60,14 +60,21 @@ Repository = "https://github.com/Onto-Med/concept-graphs" vectorstore = [ "marqo==3.18.0", ] -rag = [ +llm = [ "langchain", "langchain-core", - "langchain-community", - "langchain-text-splitters", "langchain-openai", "langchain-ollama", ] +rag = [ + { include-group = "llm" }, + "langchain-community", + "langchain-text-splitters", +] +query-expansion = [ + { include-group = "llm" }, + "pydantic>=2,<3", +] test = [ "ruff>=0.14.0", ] @@ -77,6 +84,7 @@ package = false default-groups = [ "vectorstore", "rag", + "query-expansion", ] [tool.pytest.ini_options] diff --git a/src/api/context.py b/src/api/context.py index cd18b62..65d75cf 100644 --- a/src/api/context.py +++ b/src/api/context.py @@ -16,9 +16,22 @@ class ActiveRAG: process: str ready: bool = False initializing: bool = False + error: str | None = None + + def mark_ready(self) -> None: + self.ready = True + self.initializing = False + self.error = None + + def mark_not_ready(self, error: str | None = None) -> None: + self.ready = False + self.initializing = False + self.error = error def switch_readiness(self): self.ready = not self.ready + if self.ready: + self.error = None @dataclass diff --git a/src/api/routes/rag.py b/src/api/routes/rag.py index 9be4e11..27805cb 100644 --- a/src/api/routes/rag.py +++ b/src/api/routes/rag.py @@ -20,7 +20,7 @@ ) from src.common.threads import StoppableThread from src.rag.marqo_rag_utils import extract_text_from_highlights -from src.rag.rag import RAG +from src.rag.rag import RAG, no_source_answer def create_rag_blueprint(rag, processes, storage, pipeline): @@ -74,7 +74,7 @@ def init_rag(): return jsonify("Starting initializing RAG component."), int( HTTPResponses.OK ) - rag.active_by_process[process].switch_readiness() + rag.active_by_process[process].mark_ready() return jsonify("Initialized RAG component."), int(HTTPResponses.OK) return jsonify( f"Wrong content type '{request.headers.get('Content-Type')}'; need 'application/json'" @@ -107,17 +107,25 @@ def rag_question(): if not question: return jsonify("No question supplied."), int(HTTPResponses.BAD_REQUEST) + chunks = active_rag.vectorstore.get_chunks( + question, + filter_by={"doc_id": doc_ids} if len(doc_ids) > 0 else None, + limit=doc_part_limit, + ) + if not chunks: + logging.warning( + "RAG vector store returned no chunks for process '%s'.", process + ) + return jsonify( + answer=no_source_answer(language), + info=json.dumps({}, ensure_ascii=False), + ), int(HTTPResponses.OK) + documents = list( zip( *itemgetter(1, -1)( extract_text_from_highlights( - active_rag.vectorstore.get_chunks( - question, - filter_by=( - {"doc_id": doc_ids} if len(doc_ids) > 0 else None - ), - limit=doc_part_limit, - ), + chunks, token_limit=150, lang=language, ) @@ -125,6 +133,16 @@ def rag_question(): ) ) rag_documents = active_rag.rag.documents_from(documents, concat_by="doc_id") + if not rag_documents: + logging.warning( + "RAG chunk extraction produced no source documents for process '%s'.", + process, + ) + return jsonify( + answer=no_source_answer(language), + info=json.dumps({}, ensure_ascii=False), + ), int(HTTPResponses.OK) + success, answer = active_rag.rag.build_and_invoke( question, documents=rag_documents ) diff --git a/src/api/routes/status.py b/src/api/routes/status.py index 763aa4b..9b35886 100644 --- a/src/api/routes/status.py +++ b/src/api/routes/status.py @@ -48,9 +48,33 @@ def get_rag_status(): process = string_conformity(request.args.get("process", "default")) active_rag = rag.active_by_process.get(process) if active_rag is not None: - return jsonify(active=active_rag.ready, name=process, error=None), int( - HTTPResponses.OK - ) + vectorstore_index = getattr(active_rag.vectorstore, "index_name", None) + document_count = None + try: + vectorstore_filled = active_rag.vectorstore.is_filled() + document_count_getter = getattr( + active_rag.vectorstore, "document_count", None + ) + if document_count_getter is not None: + document_count = document_count_getter() + except Exception as exc: + vectorstore_filled = False + error = f"Could not inspect RAG vector store: {exc}" + else: + error = active_rag.error + if active_rag.ready and not vectorstore_filled: + error = ( + "RAG component is marked ready, but its vector store is empty." + ) + return jsonify( + active=active_rag.ready and vectorstore_filled, + name=process, + error=error, + initializing=active_rag.initializing, + vectorstore_filled=vectorstore_filled, + vectorstore_index=vectorstore_index, + vectorstore_document_count=document_count, + ), int(HTTPResponses.OK) err_string = "The RAG component is not initialized for this process." return jsonify(active=False, name=process, error=err_string), int( HTTPResponses.NOT_FOUND diff --git a/src/api/services/rag_vectorstore.py b/src/api/services/rag_vectorstore.py index d2958b8..89ccba8 100644 --- a/src/api/services/rag_vectorstore.py +++ b/src/api/services/rag_vectorstore.py @@ -8,6 +8,7 @@ from marqo.errors import MarqoError +from src.common.parsing import string_conformity from src.pipeline.load_utils import FactoryLoader from src.pipeline.status import StepsName from src.rag.embedding_stores.base import ChunkEmbeddingStore @@ -20,8 +21,11 @@ def initialize_chunk_vectorstore( chunk_store: str = "src.rag.embedding_stores.marqo.MarqoChunkEmbeddingStore", force_init: bool = False, ): + process_name = string_conformity(process_name) if config is None: config = {"index_settings": None} + else: + config = dict(config) if config.get("index_settings", None) is None or len(config["index_settings"]) == 0: config["index_settings"] = { "type": "structured", @@ -43,19 +47,56 @@ def initialize_chunk_vectorstore( "type": "text", "features": ["lexical_search", "filter"], }, + { + "name": "chunk_index", + "type": "int", + "features": ["filter"], + }, + { + "name": "chunk_start", + "type": "int", + "features": ["filter"], + }, + { + "name": "chunk_end", + "type": "int", + "features": ["filter"], + }, {"name": "text", "type": "text", "features": ["lexical_search"]}, ], "tensorFields": ["text"], } + index_name = f"{process_name}_rag" + url = config.pop("url", "http://localhost") + port = config.pop("port", 8882) + logging.info( + "[initialize_chunk_vectorstore] Initializing RAG vector store index='%s' url='%s' port=%s force_init=%s", + index_name, + url, + port, + force_init, + ) chunk_store: ChunkEmbeddingStore = cast( ChunkEmbeddingStore, locate(chunk_store) ).from_config( - index_name=f"{process_name}_rag", - url=config.pop("url", "http://localhost"), - port=config.pop("port", 8882), + index_name=index_name, + url=url, + port=port, force_init=force_init, **config, ) + try: + logging.info( + "[initialize_chunk_vectorstore] RAG vector store index='%s' filled=%s", + index_name, + chunk_store.is_filled(), + ) + except Exception as exc: + logging.warning( + "[initialize_chunk_vectorstore] Could not inspect RAG vector store index='%s': %s", + index_name, + exc, + ) return chunk_store @@ -67,6 +108,10 @@ def fill_chunk_vectorstore(process: str, rag, storage, pipeline, **kwargs) -> bo :param kwargs: e.g. splitter=splitter-config-dict :return: """ + process = string_conformity(process) + logging.info( + "[fill_chunk_vectorstore] Starting RAG indexing for process '%s'.", process + ) _splitter_class = PreprocessedSpacyTextSplitter _split_options = { "doc_metadata_key": kwargs.get("splitter", {}).pop( @@ -91,17 +136,33 @@ def fill_chunk_vectorstore(process: str, rag, storage, pipeline, **kwargs) -> bo return False if not _rag.initializing: _rag.initializing = True + data_path = str(pathlib.Path(storage.file_storage_dir, process).resolve()) + logging.info( + "[fill_chunk_vectorstore] Loading DATA object for process '%s' from '%s'. Active pipeline keys: %s", + process, + data_path, + list(pipeline.active_objects.keys()), + ) data_obj = FactoryLoader.with_active_objects( - str(pathlib.Path(storage.file_storage_dir, process).resolve()), + data_path, process, pipeline.active_objects, StepsName.DATA, ) if data_obj is None: - logging.error( - f"[fill_chunk_vectorstore] Data object not initialized for process '{process}'. See logs for more information." + error = ( + f"Data object not initialized for process '{process}'. " + "Run/load the pipeline data step before initializing RAG." ) + logging.error("[fill_chunk_vectorstore] %s", error) + _rag.mark_not_ready(error) return False + processed_docs = getattr(data_obj, "processed_docs", None) + logging.info( + "[fill_chunk_vectorstore] Loaded DATA object type=%s processed_docs=%s", + type(data_obj).__name__, + "missing" if processed_docs is None else len(processed_docs), + ) splitter = _splitter_class(**_splitter_options) try: @@ -111,24 +172,83 @@ def fill_chunk_vectorstore(process: str, rag, storage, pipeline, **kwargs) -> bo "[fill_chunk_vectorstore] Could not reset vectorstore: %s", e ) - _documents = splitter.split_preprocessed_sentences( - data_obj.processed_docs, **_split_options + try: + documents = list( + splitter.split_preprocessed_sentences( + data_obj.processed_docs, **_split_options + ) + ) + except (AttributeError, TypeError, ValueError) as exc: + error = f"Could not split processed documents for RAG process '{process}': {exc}" + logging.error("[fill_chunk_vectorstore] %s", error) + _rag.mark_not_ready(error) + return False + + logging.info( + "[fill_chunk_vectorstore] Split process '%s' into %s document chunk groups.", + process, + len(documents), ) _field = "text" - _rag.vectorstore.add_chunks( - [ - dict( - {_field: d}, - **{k: t[1][k] for k in _split_options.get("keep_metadata", [])}, + chunks = [] + for chunk_group, metadata in documents: + kept_metadata = { + key: metadata.get(key) + for key in _split_options.get("keep_metadata", []) + if key in metadata + } + chunk_offsets = metadata.get("chunk_offsets", []) + for chunk_index, chunk in enumerate(chunk_group): + chunk_start, chunk_end = ( + chunk_offsets[chunk_index] + if chunk_index < len(chunk_offsets) + else (None, None) ) - for t in _documents - for d in t[0] - ], + chunk_metadata = {"chunk_index": chunk_index} + if chunk_start is not None: + chunk_metadata["chunk_start"] = chunk_start + if chunk_end is not None: + chunk_metadata["chunk_end"] = chunk_end + chunks.append( + dict( + {_field: chunk}, + **chunk_metadata, + **kept_metadata, + ) + ) + logging.info( + "[fill_chunk_vectorstore] Prepared %s chunks for RAG index '%s'.", + len(chunks), + getattr(_rag.vectorstore, "index_name", f"{process}_rag"), + ) + if not chunks: + error = f"No RAG chunks were produced for process '{process}'." + logging.error("[fill_chunk_vectorstore] %s", error) + _rag.mark_not_ready(error) + return False + + _rag.vectorstore.add_chunks( + chunks, # _field, ) + logging.info( + "[fill_chunk_vectorstore] Submitted %s chunks to RAG index '%s'.", + len(chunks), + getattr(_rag.vectorstore, "index_name", f"{process}_rag"), + ) - _rag.initializing = False - _rag.switch_readiness() + if not _rag.vectorstore.is_filled(): + error = f"RAG vector store for process '{process}' is still empty after filling." + logging.error("[fill_chunk_vectorstore] %s", error) + _rag.mark_not_ready(error) + return False + + logging.info( + "[fill_chunk_vectorstore] Finished RAG indexing for process '%s'. vectorstore_filled=True document_count=%s", + process, + getattr(_rag.vectorstore, "document_count", lambda: "unknown")(), + ) + _rag.mark_ready() return True else: logging.warning("[fill_chunk_vectorstore] Already initializing") diff --git a/src/pruning/unimodal.py b/src/pruning/unimodal.py index 7f3404d..17fcd52 100644 --- a/src/pruning/unimodal.py +++ b/src/pruning/unimodal.py @@ -17,8 +17,7 @@ import numpy as np from scipy.stats import binomtest -logger = logging.getLogger() -logger.setLevel("DEBUG") +logger = logging.getLogger(__name__) # clip log of binomtest p-values to this value. MAX_NEG_LOG = np.log(np.finfo(np.float64).max) @@ -152,19 +151,12 @@ def _compute_significance( ) d["significance"] = None - try: - max_sig = max( - [s for s in graph.edges(data=True) if s is not None], - key=lambda edge: edge[2].get("significance", 0.0), - default=( - 0, - 0, - {"significance": 0.0}, - ), - )[2]["significance"] - except TypeError as te: # ToDo: there were case where TypeErrors were thrown (comparison btw. 'NoneType') but I thought I made sure no 'NoneType' was allowed in the list... - logging.error(f"{te}\n-->\t{graph.edges(data=True)}") - max_sig = 0.0 + valid_significances = [ + d["significance"] + for _, _, d in graph.edges(data=True) + if d.get("significance") is not None + ] + max_sig = max(valid_significances, default=0.0) for _, _, d in graph.edges(data=True): if d["significance"] is None: d["significance"] = max_sig @@ -327,6 +319,18 @@ def _compute_significance( # raise(ValueError('The graph must be given as a, igraph.Graph, a DataFrame or a list of 3-tuples')) +def _as_binomial_count(value, name: str) -> int: + """Convert weighted graph counts to the integer counts required by binomtest.""" + if value < 0: + raise ValueError(f"{name} must be non-negative") + count = int(round(value)) + if not np.isclose(value, count): + logger.debug( + "Rounded non-integer binomial count %s=%s to %s", name, value, count + ) + return count + + def _pvalue_undirected(w, ku, kv, q): """ Compute the pvalue for the undirected edge null model. @@ -338,10 +342,15 @@ def _pvalue_undirected(w, ku, kv, q): @keyparamword q: total incident weight of all vertices divided by two. Similar to the total number of edges in the graph. """ if not all(v is not None for v in [w, ku, kv, q]): - raise ValueError + raise ValueError("binomial inputs must not be None") + if q <= 0: + raise ValueError("total edge weight must be positive") + k = _as_binomial_count(w, "w") + n = _as_binomial_count(q, "q") p = ku * kv * 1.0 / q / q / 2.0 - return binomtest(k=w, n=q, p=p, alternative="greater") + p = min(max(p, 0.0), 1.0) + return binomtest(k=k, n=n, p=p, alternative="greater").pvalue def _pvalue_directed(w_uv, ku_out, kv_in, q): @@ -355,7 +364,12 @@ def _pvalue_directed(w_uv, ku_out, kv_in, q): @param q: Total sum of all edge weights in the graph. """ if not all(v is not None for v in [w_uv, ku_out, kv_in, q]): - raise ValueError + raise ValueError("binomial inputs must not be None") + if q <= 0: + raise ValueError("total edge weight must be positive") + k = _as_binomial_count(w_uv, "w_uv") + n = _as_binomial_count(q, "q") p = 1.0 * ku_out * kv_in / q / q / 1.0 - return binomtest(k=w_uv, n=q, p=p, alternative="greater") + p = min(max(p, 0.0), 1.0) + return binomtest(k=k, n=n, p=p, alternative="greater").pvalue diff --git a/src/query_expansion/__init__.py b/src/query_expansion/__init__.py new file mode 100644 index 0000000..789fd16 --- /dev/null +++ b/src/query_expansion/__init__.py @@ -0,0 +1,35 @@ +"""LLM-driven query expansion with optional source grounding.""" + +from src.query_expansion.generator import ( + LangChainExpansionGenerator, + PydanticAIExpansionGenerator, +) +from src.query_expansion.models import ( + ExpansionGeneration, + GeneratedExpansionCandidate, + GroundedExpansionCandidate, + GroundingEvidence, + GroundingOptions, + GroundingStatus, + LLMConfig, + QueryExpansionRequest, + QueryExpansionResponse, + SourceConfig, +) +from src.query_expansion.service import QueryExpansionService + +__all__ = [ + "ExpansionGeneration", + "GeneratedExpansionCandidate", + "GroundedExpansionCandidate", + "GroundingEvidence", + "GroundingOptions", + "GroundingStatus", + "LangChainExpansionGenerator", + "LLMConfig", + "PydanticAIExpansionGenerator", + "QueryExpansionRequest", + "QueryExpansionResponse", + "QueryExpansionService", + "SourceConfig", +] diff --git a/src/query_expansion/categories.py b/src/query_expansion/categories.py new file mode 100644 index 0000000..2d0789f --- /dev/null +++ b/src/query_expansion/categories.py @@ -0,0 +1,35 @@ +"""Query-expansion category definitions.""" + +from typing import Literal + +ExpansionCategory = Literal[ + "synonym", + "medication", + "diagnosis", + "symptom", + "procedure", + "abbreviation", + "broader_term", + "narrower_term", + "related_term", +] + +DEFAULT_EXPANSION_CATEGORIES: tuple[ExpansionCategory, ...] = ( + "synonym", + "medication", + "diagnosis", + "symptom", + "procedure", +) + +CATEGORY_DESCRIPTIONS: dict[ExpansionCategory, str] = { + "synonym": "Synonyms, near-synonyms, spelling variants, or lay terms.", + "medication": "Medications or drug classes associated with the input term.", + "diagnosis": "Diagnoses or diagnostic entities associated with the input term.", + "symptom": "Symptoms, signs, or clinical findings associated with the input term.", + "procedure": "Procedures, interventions, diagnostics, or treatments associated with the input term.", + "abbreviation": "Common abbreviations or expanded forms.", + "broader_term": "Broader parent concepts.", + "narrower_term": "Narrower child concepts.", + "related_term": "Other clinically or semantically related terms.", +} diff --git a/src/query_expansion/generator.py b/src/query_expansion/generator.py new file mode 100644 index 0000000..aefca6a --- /dev/null +++ b/src/query_expansion/generator.py @@ -0,0 +1,151 @@ +"""LLM generation for query expansion.""" + +import json +import re +from collections.abc import Callable +from typing import Any, Protocol + +from src.query_expansion.categories import CATEGORY_DESCRIPTIONS +from src.query_expansion.models import ExpansionGeneration, QueryExpansionRequest + + +class ExpansionGenerator(Protocol): + """Protocol for LLM-backed expansion generators.""" + + def generate(self, request: QueryExpansionRequest) -> ExpansionGeneration: + """Generate raw, ungrounded expansion candidates for a request.""" + + +class LangChainExpansionGenerator: + """LangChain-backed structured generator. + + The generator keeps LangChain as the default project LLM framework while + still validating all LLM output with the Pydantic ``ExpansionGeneration`` + model. A concrete LangChain chat model/runnable can be injected for tests or + custom deployments. If none is provided, a small provider factory supports + ``ollama`` and OpenAI-compatible chat endpoints. + """ + + def __init__( + self, + llm: Any | None = None, + llm_factory: Callable[[QueryExpansionRequest], Any] | None = None, + ): + self._llm = llm + self._llm_factory = llm_factory + + def generate(self, request: QueryExpansionRequest) -> ExpansionGeneration: + """Generate and Pydantic-validate structured LangChain output.""" + llm = self._llm or self._build_llm(request) + prompt = build_generation_prompt(request) + + if hasattr(llm, "with_structured_output"): + structured_llm = llm.with_structured_output(ExpansionGeneration) + result = structured_llm.invoke(prompt) + return _validate_generation(result) + + result = llm.invoke(prompt) + return _validate_generation(_extract_json_payload(result)) + + def _build_llm(self, request: QueryExpansionRequest) -> Any: + if self._llm_factory is not None: + return self._llm_factory(request) + + provider = request.llm.options.get("provider", "ollama") + if provider == "ollama": + try: + from langchain_ollama import ChatOllama + except ModuleNotFoundError as exc: + raise RuntimeError( + "langchain-ollama is required for Ollama query expansion." + ) from exc + + return ChatOllama( + model=request.llm.model, + base_url=request.llm.options.get("base_url", "http://localhost:11434"), + temperature=request.llm.options.get("temperature", 0.0), + ) + + if provider in {"openai", "blablador"}: + try: + from langchain_openai import ChatOpenAI + except ModuleNotFoundError as exc: + raise RuntimeError( + "langchain-openai is required for OpenAI-compatible query expansion." + ) from exc + + return ChatOpenAI( + model=request.llm.model, + base_url=request.llm.options.get("base_url"), + api_key=request.llm.options.get("api_key"), + temperature=request.llm.options.get("temperature", 0.0), + ) + + raise ValueError(f"Unsupported LangChain query-expansion provider: {provider}") + + +class PydanticAIExpansionGenerator: + """PydanticAI-backed structured generator. + + The import is intentionally lazy so the rest of the query-expansion package can + be imported in environments where pydantic-ai is not installed yet. + """ + + def generate(self, request: QueryExpansionRequest) -> ExpansionGeneration: + """Run a PydanticAI agent and return structured expansion candidates.""" + try: + from pydantic_ai import Agent + except ModuleNotFoundError as exc: + raise RuntimeError( + "pydantic-ai is required for LLM query expansion. Install it or use " + "a test/fake ExpansionGenerator implementation." + ) from exc + + prompt = build_generation_prompt(request) + agent = Agent( + request.llm.model, + result_type=ExpansionGeneration, + system_prompt=request.llm.system_prompt, + **request.llm.options, + ) + result = agent.run_sync(prompt) + return result.data + + +def _validate_generation(value: Any) -> ExpansionGeneration: + if isinstance(value, ExpansionGeneration): + return value + return ExpansionGeneration.model_validate(value) + + +def _extract_json_payload(value: Any) -> dict[str, Any]: + content = getattr(value, "content", value) + if isinstance(content, dict): + return content + if isinstance(content, str): + text = content.strip() + fenced_json = re.search(r"```(?:json)?\s*(.*?)```", text, re.DOTALL) + if fenced_json: + text = fenced_json.group(1).strip() + return json.loads(text) + raise TypeError(f"Cannot extract JSON query-expansion payload from {type(value)!r}") + + +def build_generation_prompt(request: QueryExpansionRequest) -> str: + """Build the prompt used by the LLM expansion generator.""" + category_descriptions = { + category: CATEGORY_DESCRIPTIONS.get(category, category) + for category in request.categories + } + return ( + "Generate query-expansion candidates for the provided term. " + "Return only candidates that are useful for search/query expansion. " + "Group each candidate by one of the requested categories.\n\n" + f"Term: {request.term}\n" + f"Language: {request.language}\n" + f"Limit per category: {request.limit_per_category}\n" + f"Categories: {json.dumps(category_descriptions, ensure_ascii=False)}\n\n" + "Return JSON matching this schema exactly: " + '{"candidates": [{"term": "...", "category": "...", "rationale": "..."}]}. ' + "The output is validated as an ExpansionGeneration Pydantic model." + ) diff --git a/src/query_expansion/grounding.py b/src/query_expansion/grounding.py new file mode 100644 index 0000000..394315e --- /dev/null +++ b/src/query_expansion/grounding.py @@ -0,0 +1,55 @@ +"""Grounding generated query-expansion candidates against sources.""" + +from collections import defaultdict + +from src.query_expansion.models import ( + GeneratedExpansionCandidate, + GroundedExpansionCandidate, + GroundingOptions, + GroundingStatus, +) +from src.query_expansion.sources.base import ExpansionSource + + +def ground_candidate( + candidate: GeneratedExpansionCandidate, + sources: list[ExpansionSource], + options: GroundingOptions, +) -> GroundedExpansionCandidate | None: + """Ground one LLM-generated candidate against all configured sources. + + Returns ``None`` when the candidate should be filtered out according to the + grounding options. + """ + evidence = [item for source in sources for item in source.ground(candidate)] + confidence = max((item.score for item in evidence), default=0.0) + if confidence >= options.minimum_score and evidence: + status = GroundingStatus.GROUNDED + elif options.reject_below_minimum and confidence < options.minimum_score: + status = GroundingStatus.REJECTED + else: + status = GroundingStatus.LLM_ONLY + + if status == GroundingStatus.LLM_ONLY and not options.include_llm_only: + return None + if status == GroundingStatus.REJECTED: + return None + + return GroundedExpansionCandidate( + term=candidate.term, + category=candidate.category, + status=status, + confidence=confidence, + evidence=evidence, + rationale=candidate.rationale, + ) + + +def group_grounded_candidates( + candidates: list[GroundedExpansionCandidate], +) -> dict[str, list[GroundedExpansionCandidate]]: + """Group final candidates by semantic expansion category.""" + grouped = defaultdict(list) + for candidate in candidates: + grouped[candidate.category].append(candidate) + return dict(grouped) diff --git a/src/query_expansion/models.py b/src/query_expansion/models.py new file mode 100644 index 0000000..c24f94e --- /dev/null +++ b/src/query_expansion/models.py @@ -0,0 +1,234 @@ +"""Pydantic data models for query expansion. + +The query-expansion workflow is intentionally modeled in three stages: + +1. A request describes the input term, requested semantic categories, LLM + settings, and optional grounding sources. +2. The LLM returns an ``ExpansionGeneration`` containing raw generated + candidates. +3. The service grounds those candidates against terminology/ontology sources and + returns a ``QueryExpansionResponse`` with evidence and grounding status. + +The models are Pydantic/PydanticAI-friendly: ``ExpansionGeneration`` can be used +as a structured output model for a PydanticAI agent while the rest of the package +remains importable without PydanticAI installed. +""" + +from enum import StrEnum +from typing import Any + +from pydantic import BaseModel, Field, field_validator + +from src.query_expansion.categories import ( + DEFAULT_EXPANSION_CATEGORIES, + ExpansionCategory, +) + + +class GroundingStatus(StrEnum): + """How well an LLM-generated candidate is supported by external sources.""" + + GROUNDED = "grounded" + """The candidate is supported by at least one source above the score threshold.""" + + PARTIALLY_GROUNDED = "partially_grounded" + """Reserved for future source adapters that can provide weaker relation evidence.""" + + LLM_ONLY = "llm_only" + """The candidate was generated by the LLM but not found in configured sources.""" + + REJECTED = "rejected" + """The candidate failed grounding criteria and is omitted from normal responses.""" + + +class SourceConfig(BaseModel): + """Configuration for a grounding source. + + A source is an ontology, terminology, dictionary, or API that can validate or + enrich LLM-generated candidates. Currently only ``type="local"`` is + implemented. + """ + + name: str = Field(description="Human-readable source name used in evidence.") + type: str = Field(default="local", description="Source adapter type.") + path: str | None = Field( + default=None, description="Path for local YAML/JSON terminology sources." + ) + url: str | None = Field( + default=None, description="Base URL for future HTTP/API-backed sources." + ) + options: dict[str, Any] = Field( + default_factory=dict, description="Adapter-specific source options." + ) + + +class GroundingOptions(BaseModel): + """Options controlling how generated candidates are filtered after grounding.""" + + include_llm_only: bool = Field( + default=True, + description="Include candidates that are not grounded in any configured source.", + ) + minimum_score: float = Field( + default=0.0, + ge=0.0, + le=1.0, + description="Minimum grounding score required to treat a candidate as grounded.", + ) + reject_below_minimum: bool = Field( + default=False, + description="Drop candidates below minimum_score instead of returning them as llm_only.", + ) + + +class LLMConfig(BaseModel): + """Configuration for the LLM used as primary expansion generator. + + The ``model`` value is passed to the selected LLM generator. Provider-specific + parameters can be passed through ``options``. For the LangChain generator, + ``options.provider`` currently supports ``ollama``, ``openai``, and + ``blablador``. + """ + + model: str = Field(description="LLM model identifier for the selected generator.") + system_prompt: str | None = Field( + default=None, description="Optional system prompt for the expansion agent." + ) + instructions: str | None = Field( + default=None, + description="Reserved for additional user/deployment instructions.", + ) + options: dict[str, Any] = Field( + default_factory=dict, + description="Additional provider/generator-specific keyword arguments.", + ) + + +class QueryExpansionRequest(BaseModel): + """Input request for LLM-first, source-grounded query expansion.""" + + term: str = Field( + min_length=1, description="Input term to expand, e.g. 'myocardial infarction'." + ) + language: str = Field( + default="en", description="Language code used for prompting and source lookup." + ) + categories: list[ExpansionCategory] = Field( + default_factory=lambda: list(DEFAULT_EXPANSION_CATEGORIES), + description="Semantic expansion categories requested from the LLM.", + ) + limit_per_category: int = Field( + default=10, + ge=1, + le=100, + description="Maximum number of generated candidates per category.", + ) + llm: LLMConfig = Field(description="Required LLM generator configuration.") + sources: list[SourceConfig] = Field( + default_factory=list, + description="Optional sources used to ground/validate generated candidates.", + ) + grounding: GroundingOptions = Field( + default_factory=GroundingOptions, + description="Grounding and post-filtering behavior.", + ) + + @field_validator("term") + @classmethod + def strip_term(cls, value: str) -> str: + """Normalize accidental surrounding whitespace in the input term.""" + return value.strip() + + +class GeneratedExpansionCandidate(BaseModel): + """Raw candidate produced by the LLM before grounding. + + ``rationale`` is the LLM's short explanation for why the candidate belongs to + the category. It is not used as proof. It is carried through to the response + to support debugging, review, or UI display. + """ + + term: str = Field( + min_length=1, + description="Generated expansion term, e.g. 'heart attack' or 'aspirin'.", + ) + category: ExpansionCategory = Field( + description="Semantic category assigned by the LLM." + ) + rationale: str | None = Field( + default=None, + description="Optional LLM explanation for the candidate/category choice; not grounding evidence.", + ) + + @field_validator("term") + @classmethod + def strip_candidate_term(cls, value: str) -> str: + """Normalize accidental surrounding whitespace in generated candidates.""" + return value.strip() + + +class ExpansionGeneration(BaseModel): + """Structured LLM output model for PydanticAI. + + A PydanticAI agent should return this object. The service will then ground + each candidate against configured sources. + """ + + candidates: list[GeneratedExpansionCandidate] = Field( + default_factory=list, description="LLM-generated expansion candidates." + ) + + +class GroundingEvidence(BaseModel): + """Evidence that a source supports an expansion candidate.""" + + source: str = Field(description="Name of the source that produced the evidence.") + matched_term: str = Field( + description="Source term/synonym that matched the generated candidate." + ) + score: float = Field( + ge=0.0, + le=1.0, + description="Grounding confidence score assigned by the source adapter.", + ) + relation: str | None = Field( + default=None, + description="Source relation type, e.g. exact_or_synonym, broader, narrower.", + ) + source_id: str | None = Field( + default=None, description="Optional concept/code identifier from the source." + ) + metadata: dict[str, Any] = Field( + default_factory=dict, + description="Additional source-specific evidence metadata.", + ) + + +class GroundedExpansionCandidate(BaseModel): + """Final expansion candidate after grounding and filtering.""" + + term: str = Field(description="Expansion term returned to the caller.") + category: ExpansionCategory = Field(description="Semantic expansion category.") + status: GroundingStatus = Field(description="Grounding classification.") + confidence: float = Field( + ge=0.0, + le=1.0, + description="Highest grounding evidence score, or 0 for ungrounded LLM-only terms.", + ) + evidence: list[GroundingEvidence] = Field( + default_factory=list, description="Source evidence supporting the candidate." + ) + rationale: str | None = Field( + default=None, + description="Optional LLM rationale copied from the generated candidate.", + ) + + +class QueryExpansionResponse(BaseModel): + """Response returned by the query-expansion service/API.""" + + term: str = Field(description="Original input term.") + language: str = Field(description="Request language.") + expansions: dict[ExpansionCategory, list[GroundedExpansionCandidate]] = Field( + description="Expansion candidates grouped by requested category." + ) diff --git a/src/query_expansion/service.py b/src/query_expansion/service.py new file mode 100644 index 0000000..f651b76 --- /dev/null +++ b/src/query_expansion/service.py @@ -0,0 +1,84 @@ +"""Query-expansion orchestration service.""" + +from src.query_expansion.generator import ( + ExpansionGenerator, + PydanticAIExpansionGenerator, +) +from src.query_expansion.grounding import ground_candidate, group_grounded_candidates +from src.query_expansion.models import ( + QueryExpansionRequest, + QueryExpansionResponse, + SourceConfig, +) +from src.query_expansion.sources.base import ExpansionSource +from src.query_expansion.sources.local import LocalTerminologySource + + +def source_from_config(config: SourceConfig) -> ExpansionSource: + """Create a grounding source adapter from request configuration.""" + if config.type == "local": + if config.path is None: + raise ValueError( + f"Local query-expansion source '{config.name}' needs a path." + ) + return LocalTerminologySource(config.name, config.path) + raise NotImplementedError( + f"Query-expansion source type '{config.type}' is not implemented yet." + ) + + +class QueryExpansionService: + """Coordinate LLM generation and optional source grounding. + + The service is intentionally independent from Flask so it can be used from an + API route, a CLI, tests, or future batch jobs. By default it uses the + PydanticAI-backed generator, but tests or deployments can inject any object + implementing the ``ExpansionGenerator`` protocol. + """ + + def __init__(self, generator: ExpansionGenerator | None = None): + """Create a service with either a custom or default LLM generator.""" + self.generator = generator or PydanticAIExpansionGenerator() + + def expand( + self, + request: QueryExpansionRequest, + sources: list[ExpansionSource] | None = None, + ) -> QueryExpansionResponse: + """Generate candidates, ground them, and group them by category. + + Args: + request: User request containing the term, categories, LLM config, + source config, and grounding options. + sources: Optional pre-built source adapters. If omitted, adapters are + created from ``request.sources``. + + Returns: + A response containing grounded and/or LLM-only candidates grouped by + category. + """ + sources = ( + [source_from_config(source_config) for source_config in request.sources] + if sources is None + else sources + ) + generated = self.generator.generate(request) + grounded = [ + grounded_candidate + for candidate in generated.candidates + if candidate.category in request.categories + if ( + grounded_candidate := ground_candidate( + candidate, sources, request.grounding + ) + ) + is not None + ] + grouped = group_grounded_candidates(grounded) + return QueryExpansionResponse( + term=request.term, + language=request.language, + expansions={ + category: grouped.get(category, []) for category in request.categories + }, + ) diff --git a/src/query_expansion/sources/__init__.py b/src/query_expansion/sources/__init__.py new file mode 100644 index 0000000..8fc4a73 --- /dev/null +++ b/src/query_expansion/sources/__init__.py @@ -0,0 +1,11 @@ +"""Query-expansion source adapters.""" + +from src.query_expansion.sources.base import ExpansionSource +from src.query_expansion.sources.http import HTTPExpansionSource +from src.query_expansion.sources.local import LocalTerminologySource + +__all__ = [ + "ExpansionSource", + "HTTPExpansionSource", + "LocalTerminologySource", +] diff --git a/src/query_expansion/sources/base.py b/src/query_expansion/sources/base.py new file mode 100644 index 0000000..05b73e0 --- /dev/null +++ b/src/query_expansion/sources/base.py @@ -0,0 +1,16 @@ +"""Source interfaces for grounding query-expansion candidates.""" + +from abc import ABC, abstractmethod + +from src.query_expansion.models import GeneratedExpansionCandidate, GroundingEvidence + + +class ExpansionSource(ABC): + """Base class for terminology/ontology/source adapters.""" + + name: str + + @abstractmethod + def ground(self, candidate: GeneratedExpansionCandidate) -> list[GroundingEvidence]: + """Return grounding evidence for a generated candidate.""" + raise NotImplementedError diff --git a/src/query_expansion/sources/http.py b/src/query_expansion/sources/http.py new file mode 100644 index 0000000..c0fadd9 --- /dev/null +++ b/src/query_expansion/sources/http.py @@ -0,0 +1,20 @@ +"""Placeholder HTTP source adapter. + +Network-backed terminology adapters should subclass ``ExpansionSource`` and +implement source-specific authentication, lookup, and response mapping. +""" + +from src.query_expansion.models import GeneratedExpansionCandidate, GroundingEvidence +from src.query_expansion.sources.base import ExpansionSource + + +class HTTPExpansionSource(ExpansionSource): + def __init__(self, name: str, url: str, options: dict | None = None): + self.name = name + self.url = url + self.options = options or {} + + def ground(self, candidate: GeneratedExpansionCandidate) -> list[GroundingEvidence]: + raise NotImplementedError( + "HTTP query-expansion sources are not implemented yet." + ) diff --git a/src/query_expansion/sources/local.py b/src/query_expansion/sources/local.py new file mode 100644 index 0000000..b6aa8c8 --- /dev/null +++ b/src/query_expansion/sources/local.py @@ -0,0 +1,86 @@ +"""Local file based grounding sources.""" + +import json +from pathlib import Path +from typing import Any + +import yaml + +from src.query_expansion.models import GeneratedExpansionCandidate, GroundingEvidence +from src.query_expansion.sources.base import ExpansionSource + + +class LocalTerminologySource(ExpansionSource): + """Ground candidates against a small local terminology file. + + Supported file shapes: + + ```yaml + terms: + - id: C001 + term: myocardial infarction + synonyms: [heart attack, MI] + ``` + + or directly: + + ```yaml + myocardial infarction: + synonyms: [heart attack, MI] + ``` + """ + + def __init__(self, name: str, path: str | Path): + """Load a local terminology source from YAML or JSON.""" + self.name = name + self.path = Path(path) + self.entries = self._load_entries(self.path) + + @staticmethod + def _normalize(value: str) -> str: + return " ".join(value.lower().split()) + + @classmethod + def _load_entries(cls, path: Path) -> list[dict[str, Any]]: + if path.suffix.lower() == ".json": + data = json.loads(path.read_text()) + else: + data = yaml.safe_load(path.read_text()) + if data is None: + return [] + if isinstance(data, dict) and isinstance(data.get("terms"), list): + return data["terms"] + if isinstance(data, dict): + return [ + dict({"term": term}, **(payload or {})) + for term, payload in data.items() + if isinstance(payload, dict) + ] + if isinstance(data, list): + return [entry for entry in data if isinstance(entry, dict)] + return [] + + def ground(self, candidate: GeneratedExpansionCandidate) -> list[GroundingEvidence]: + """Return exact/synonym matches for a generated candidate.""" + candidate_term = self._normalize(candidate.term) + evidence = [] + for entry in self.entries: + terms = [entry.get("term", "")] + terms.extend(entry.get("synonyms", []) or []) + normalized_terms = {self._normalize(term): term for term in terms if term} + if candidate_term in normalized_terms: + evidence.append( + GroundingEvidence( + source=self.name, + matched_term=normalized_terms[candidate_term], + score=1.0, + relation="exact_or_synonym", + source_id=entry.get("id"), + metadata={ + k: v + for k, v in entry.items() + if k not in {"id", "term", "synonyms"} + }, + ) + ) + return evidence diff --git a/src/rag/embedding_stores/marqo.py b/src/rag/embedding_stores/marqo.py index a0504ef..d311ade 100644 --- a/src/rag/embedding_stores/marqo.py +++ b/src/rag/embedding_stores/marqo.py @@ -18,6 +18,13 @@ def __init__( self._index_name: str = index_name self._config: dict = config if config is not None else {} + @property + def index_name(self) -> str: + return self._index_name + + def document_count(self) -> int: + return self._client.index(self._index_name).get_stats()["numberOfDocuments"] + def _init_index(self): self._client.create_index( index_name=self._index_name, settings_dict=self._config @@ -51,7 +58,7 @@ def from_config( return _store def is_filled(self) -> bool: - return self._client.index(self._index_name).get_stats()["numberOfDocuments"] > 0 + return self.document_count() > 0 def reset_index(self, with_settings: dict[str, Any] | None = None) -> None: _settings = ( diff --git a/src/rag/marqo_rag_utils.py b/src/rag/marqo_rag_utils.py index 6ab965e..16523c7 100644 --- a/src/rag/marqo_rag_utils.py +++ b/src/rag/marqo_rag_utils.py @@ -53,7 +53,11 @@ def truncate_text(text, token_limit, highlight=None, lang: str = "en"): """ truncates text to a token limit centered on the highlight text """ + return truncate_text_with_offsets(text, token_limit, highlight, lang)[0] + +def truncate_text_with_offsets(text, token_limit, highlight=None, lang: str = "en"): + """Truncate text and return ``(snippet, start, end)`` within the input text.""" if highlight is None: method: _method_literal = "start" center_ind = 0 # this will not be used for this start method @@ -68,7 +72,7 @@ def truncate_text(text, token_limit, highlight=None, lang: str = "en"): logging.warning( "Could not find highlight index in text; using text value. Might exceed token limit." ) - return text + return text, 0, len(text) # get the center of the highlight in chars center_ind = (max(inds) - min(inds)) // 2 + min(inds) # now map this to tokens and get the left/right char indices to achieve token limit @@ -76,9 +80,11 @@ def truncate_text(text, token_limit, highlight=None, lang: str = "en"): ind_left, ind_right = get_token_indices( text, token_limit, method=method, offset=center_ind, lang=lang ) - trunc_text = text[min(ind_left) : max(ind_right)] + snippet_start = min(ind_left) + snippet_end = max(ind_right) + trunc_text = text[snippet_start:snippet_end] - return trunc_text + return trunc_text, snippet_start, snippet_end def get_token_indices( @@ -158,15 +164,53 @@ def extract_text_from_highlights( highlight_list = hit[ResultsFields.highlights] highlight_key = list(highlight_list[0].keys())[0] highlight_text = list(highlight_list[0].values())[0] - text = hit.pop(highlight_key) + text = hit.get(highlight_key, "") + snippet_start_in_chunk = 0 + snippet_end_in_chunk = len(text) if truncate: text = " ".join(text.split()) highlight_text = " ".join(highlight_text.split()) - text = truncate_text(text, token_limit, highlight_text, lang) + text, snippet_start_in_chunk, snippet_end_in_chunk = ( + truncate_text_with_offsets(text, token_limit, highlight_text, lang) + ) + + highlight_offsets = find_highlight_index_in_text(text, highlight_text) + chunk_start = hit.get("chunk_start") + hit_metadata = { + k: hit.get(k) + for k in hit.keys() + if not k.startswith("_") and k != highlight_key + } + hit_metadata.update( + { + "retrieved_snippet": text, + "retrieved_snippet_start": snippet_start_in_chunk, + "retrieved_snippet_end": snippet_end_in_chunk, + "highlight": highlight_text, + "highlight_field": highlight_key, + "highlight_start": None + if highlight_offsets is None + else highlight_offsets[0], + "highlight_end": None + if highlight_offsets is None + else highlight_offsets[1], + "offset_unit": "retrieved_snippet_char", + "document_highlight_start": None + if highlight_offsets is None or chunk_start is None + else chunk_start + snippet_start_in_chunk + highlight_offsets[0], + "document_highlight_end": None + if highlight_offsets is None or chunk_start is None + else chunk_start + snippet_start_in_chunk + highlight_offsets[1], + "document_offset_unit": "document_char" + if chunk_start is not None + else None, + "retrieved_snippet_index": ind, + } + ) texts.append(text) highlights.append(highlight_text) - metadata.append({k: hit.get(k) for k in hit.keys() if not k.startswith("_")}) + metadata.append(hit_metadata) return highlights, texts, metadata diff --git a/src/rag/rag.py b/src/rag/rag.py index 03cb55e..d01eaa2 100644 --- a/src/rag/rag.py +++ b/src/rag/rag.py @@ -14,6 +14,27 @@ from src.rag.marqo_rag_utils import extract_text_from_highlights +def no_source_answer(language: str | None = None) -> str: + """Return a deterministic answer for RAG requests without source documents.""" + if language == "de": + return "Keine Quelle die ich finden kann." + return "No source I can find." + + +def _clean_answer(answer: Any) -> Any: + """Remove common prompt-echo tails while keeping the generated answer.""" + content = getattr(answer, "content", answer) + if not isinstance(content, str): + return answer + + text = content.strip() + for marker in ["=========", "\nFRAGE:", "\nQUESTION:", "\nQUELLEN:", "\nSOURCES:"]: + marker_index = text.find(marker) + if marker_index > 0: + return text[:marker_index].strip() + return text + + class RAG: def __init__(self, chatter: Chatter | str, language: str | None = None): self._language = language @@ -108,8 +129,9 @@ def with_prompt( """ _templates = { "en": """ - Given the following extracted parts of several different documents ("SOURCES") and a question ("QUESTION"), create a final answer one paragraph long. - Don't try to make up an answer and use the text in the SOURCES only for the answer. If you don't know the answer, just say that you don't know. + Given the following extracted parts of several different documents ("SOURCES") and a question ("QUESTION"), create a final answer one paragraph long. + Use only the SOURCES. Do not make up an answer. If you don't know the answer, say that you don't know. + Output only the final answer. Do not repeat the question, sources, separators, or instructions. Do not include analysis. QUESTION: {question} ========= SOURCES: @@ -119,8 +141,9 @@ def with_prompt( """, "de": """ Gegeben sind die folgenden Teile verschiedener Dokumente ("QUELLEN") und eine Frage ("FRAGE"), erstelle eine kurze abschließende "ANTWORT" mit etwa einer Länge eines Absatzes. - Dabei sollen die "QUELLEN" individuell betrachtet werden! Referenziere in der Antwort die "QUELLEN"! Erwähne nur positive Antworten! - Versuche niemals eine Antwort zu erfinden! Benutze ausschließlich die Texte aus den "QUELLEN" für die "ANTWORT". Wenn du keine "ANTWORT" hast, sage einfach, dass du es nicht weißt! + Betrachte die "QUELLEN" individuell. Referenziere in der Antwort die "QUELLEN". Erwähne nur positive Antworten. + Erfinde niemals eine Antwort. Benutze ausschließlich die Texte aus den "QUELLEN" für die "ANTWORT". Wenn du keine "ANTWORT" hast, sage einfach, dass du es nicht weißt. + Gib ausschließlich die finale Antwort aus. Wiederhole nicht die Frage, Quellen, Trennzeichen oder Anweisungen. Gib keine Analyse aus. FRAGE: {question} ========= QUELLEN: @@ -184,15 +207,28 @@ def with_documents( def build(self) -> Runnable: return self._prompt | self._initialized_chatter + def no_source_answer(self) -> str: + """Return a deterministic answer for missing retrieved documents.""" + return no_source_answer(self.language) + def build_and_invoke( self, question: str, documents: dict[str, Document] | None = None ): documents = self.documents if documents is None else documents + if not documents: + logging.warning( + "No RAG source documents available; skipping LLM invocation." + ) + return True, self.no_source_answer() + summaries = "\n\n".join( + document.page_content for document in documents.values() + ) try: - return True, (self._prompt | self._initialized_chatter).invoke( - {"summaries": documents.values(), "question": question}, + answer = (self._prompt | self._initialized_chatter).invoke( + {"summaries": summaries, "question": question}, return_only_outputs=True, ) + return True, _clean_answer(answer) except (LangChainException, RuntimeError, ValueError, TypeError) as e: logging.warning("RAG invocation failed: %s", e) return False, e diff --git a/src/rag/text_splitters.py b/src/rag/text_splitters.py index 152810c..eb3dc06 100644 --- a/src/rag/text_splitters.py +++ b/src/rag/text_splitters.py @@ -1,5 +1,6 @@ import logging from collections.abc import Generator, Iterable +from dataclasses import dataclass from typing import Any from langchain_text_splitters import TextSplitter @@ -8,6 +9,13 @@ logger = logging.getLogger(__name__) +@dataclass(frozen=True) +class TextSplitWithOffsets: + text: str + start: int | None + end: int | None + + class PreprocessedSpacyTextSplitter(TextSplitter): """Splitting text from preprocessed Spacy data. @@ -56,29 +64,24 @@ def split_preprocessed_sentences( ) -> Generator[list[str] | tuple[list[str], dict], None, None]: docs = [] current_doc_id = None - meta_data = {} warned_once = False + + def yield_current_docs(): + if not docs: + return None + splits_with_offsets = [self._split_with_offsets(doc) for doc in docs] + chunks, chunk_offsets = self._merge_splits_with_offsets( + splits_with_offsets, self._separator + ) + if keep_metadata is None: + return chunks + meta_data = self._metadata_from_doc(docs[0], keep_metadata) + meta_data["chunk_offsets"] = chunk_offsets + return chunks, meta_data + for sentence in sentences: - if doc_id := getattr(getattr(sentence, "_", {}), doc_metadata_key, None): - if len(meta_data) == 0 and keep_metadata is not None: - for k in getattr(sentence, "_", {}).__dict__.get("_extensions", {}): - if k in keep_metadata: - meta_data[k] = getattr(getattr(sentence, "_"), k) - splits = ( - s.text if self._strip_whitespace else s.text_with_ws for s in docs - ) - if current_doc_id is None: - current_doc_id = doc_id - if doc_id != current_doc_id: - if keep_metadata is None: - yield self._merge_splits(splits, self._separator) - else: - yield self._merge_splits(splits, self._separator), meta_data - docs = [] - meta_data = {} - current_doc_id = doc_id - docs.append(sentence) - else: + doc_id = getattr(getattr(sentence, "_", {}), doc_metadata_key, None) + if not doc_id: if not warned_once: logging.warning( f"There seems to be no metadata for '{doc_metadata_key}'" @@ -86,11 +89,90 @@ def split_preprocessed_sentences( ) warned_once = True docs.append(sentence) - if warned_once and len(docs) > 0: - splits = ( - s.text if self._strip_whitespace else s.text_with_ws for s in docs + continue + + if current_doc_id is None: + current_doc_id = doc_id + if doc_id != current_doc_id: + result = yield_current_docs() + if result is not None: + yield result + docs = [] + current_doc_id = doc_id + docs.append(sentence) + + result = yield_current_docs() + if result is not None: + yield result + + def _split_with_offsets(self, doc: Doc) -> TextSplitWithOffsets: + text = doc.text if self._strip_whitespace else doc.text_with_ws + start = getattr(getattr(doc, "_", {}), "offset_in_doc", None) + end = None if start is None else start + len(text) + return TextSplitWithOffsets(text=text, start=start, end=end) + + def _metadata_from_doc(self, doc: Doc, keep_metadata: list[str]) -> dict: + metadata = {} + if hasattr(doc, "_"): + for key in getattr(doc, "_").__dict__.get("_extensions", {}): + if key in keep_metadata: + metadata[key] = getattr(getattr(doc, "_"), key) + return metadata + + def _merge_splits_with_offsets( + self, splits: list[TextSplitWithOffsets], separator: str + ) -> tuple[list[str], list[tuple[int | None, int | None]]]: + separator_len = self._length_function(separator) + chunks = [] + offsets = [] + current_doc: list[TextSplitWithOffsets] = [] + total = 0 + + def append_current_chunk() -> None: + chunk = self._join_docs([split.text for split in current_doc], separator) + if chunk is None: + return + starts = [split.start for split in current_doc if split.start is not None] + ends = [split.end for split in current_doc if split.end is not None] + chunks.append(chunk) + offsets.append( + ( + min(starts) if starts else None, + max(ends) if ends else None, + ) ) - yield self._merge_splits(splits, self._separator) + + for split in splits: + split_len = self._length_function(split.text) + if ( + total + split_len + (separator_len if len(current_doc) > 0 else 0) + > self._chunk_size + ): + if total > self._chunk_size: + logger.warning( + "Created a chunk of size %s, which is longer than the specified %s", + total, + self._chunk_size, + ) + if len(current_doc) > 0: + append_current_chunk() + while total > self._chunk_overlap or ( + total + + split_len + + (separator_len if len(current_doc) > 0 else 0) + > self._chunk_size + and total > 0 + ): + total -= self._length_function(current_doc[0].text) + ( + separator_len if len(current_doc) > 1 else 0 + ) + current_doc = current_doc[1:] + current_doc.append(split) + total += split_len + (separator_len if len(current_doc) > 1 else 0) + + if current_doc: + append_current_chunk() + return chunks, offsets if __name__ == "__main__": diff --git a/test/api/routes/test_rag_routes.py b/test/api/routes/test_rag_routes.py index 2a4613e..6cc643b 100644 --- a/test/api/routes/test_rag_routes.py +++ b/test/api/routes/test_rag_routes.py @@ -5,9 +5,10 @@ class FakeVectorStore: - def __init__(self, process="default", filled=True): + def __init__(self, process="default", filled=True, chunks=None): self.process = process self.filled = filled + self.chunks = chunks self.chunk_requests = [] def is_filled(self): @@ -17,6 +18,8 @@ def get_chunks(self, question, filter_by=None, limit=10): self.chunk_requests.append( {"question": question, "filter_by": filter_by, "limit": limit} ) + if self.chunks is not None: + return self.chunks return [{"text": f"chunk for {self.process}"}] @@ -96,12 +99,20 @@ def test_rag_init_keeps_active_rag_per_process(monkeypatch, tmp_path): assert client.get("/status/rag?process=corpus_a").json == { "active": True, "error": None, + "initializing": False, "name": "corpus_a", + "vectorstore_document_count": None, + "vectorstore_filled": True, + "vectorstore_index": None, } assert client.get("/status/rag?process=corpus_b").json == { "active": True, "error": None, + "initializing": False, "name": "corpus_b", + "vectorstore_document_count": None, + "vectorstore_filled": True, + "vectorstore_index": None, } @@ -147,6 +158,25 @@ def test_rag_question_uses_selected_process_without_mutating_rag(monkeypatch, tm assert rag_b.documents is None +def test_rag_question_returns_no_source_answer_when_retrieval_is_empty(tmp_path): + app = _app(tmp_path) + rag_context = app.extensions["concept_graphs_context"].rag + rag_instance = FakeRAG(language="de", answer="should-not-be-called") + vector_store = FakeVectorStore(process="corpus", chunks=[]) + rag_context.active_by_process["corpus"] = ActiveRAG( + rag=rag_instance, vectorstore=vector_store, process="corpus", ready=True + ) + + response = app.test_client().get("/rag/question?process=corpus&q=question") + + assert response.status_code == 200 + assert response.json == { + "answer": "Keine Quelle die ich finden kann.", + "info": "{}", + } + assert rag_instance.invocations == [] + + def test_rag_question_returns_not_found_for_uninitialized_process(tmp_path): response = ( _app(tmp_path).test_client().get("/rag/question?process=missing&q=question") diff --git a/test/query_expansion/test_service.py b/test/query_expansion/test_service.py new file mode 100644 index 0000000..6aeae02 --- /dev/null +++ b/test/query_expansion/test_service.py @@ -0,0 +1,124 @@ +from src.query_expansion.generator import LangChainExpansionGenerator +from src.query_expansion.models import ( + ExpansionGeneration, + GeneratedExpansionCandidate, + GroundingOptions, + GroundingStatus, + LLMConfig, + QueryExpansionRequest, +) +from src.query_expansion.service import QueryExpansionService +from src.query_expansion.sources.local import LocalTerminologySource + + +class FakeStructuredLLM: + def with_structured_output(self, schema): + self.schema = schema + return self + + def invoke(self, prompt): + return { + "candidates": [ + { + "term": "heart attack", + "category": "synonym", + "rationale": "common lay term", + } + ] + } + + +class FakeJsonLLM: + def invoke(self, prompt): + return '{"candidates": [{"term": "aspirin", "category": "medication"}]}' + + +class FakeGenerator: + def generate(self, request): + return ExpansionGeneration( + candidates=[ + GeneratedExpansionCandidate( + term="heart attack", category="synonym", rationale="lay term" + ), + GeneratedExpansionCandidate( + term="aspirin", category="medication", rationale="common therapy" + ), + GeneratedExpansionCandidate( + term="unverified", category="symptom", rationale="test" + ), + ] + ) + + +def test_langchain_generator_validates_structured_output_with_pydantic(): + request = QueryExpansionRequest( + term="myocardial infarction", + categories=["synonym"], + llm=LLMConfig(model="test-model"), + ) + + generation = LangChainExpansionGenerator(llm=FakeStructuredLLM()).generate(request) + + assert generation.candidates[0].term == "heart attack" + assert generation.candidates[0].category == "synonym" + + +def test_langchain_generator_can_parse_json_fallback_output(): + request = QueryExpansionRequest( + term="myocardial infarction", + categories=["medication"], + llm=LLMConfig(model="test-model"), + ) + + generation = LangChainExpansionGenerator(llm=FakeJsonLLM()).generate(request) + + assert generation.candidates[0].term == "aspirin" + assert generation.candidates[0].category == "medication" + + +def test_query_expansion_service_generates_and_grounds_candidates(tmp_path): + source_file = tmp_path / "terms.yaml" + source_file.write_text( + """ +terms: + - id: C001 + term: myocardial infarction + synonyms: [heart attack, MI] + - id: C002 + term: aspirin +""" + ) + source = LocalTerminologySource("local", source_file) + request = QueryExpansionRequest( + term="myocardial infarction", + categories=["synonym", "medication", "symptom"], + llm=LLMConfig(model="test-model"), + ) + + response = QueryExpansionService(generator=FakeGenerator()).expand( + request, sources=[source] + ) + + assert response.term == "myocardial infarction" + assert response.expansions["synonym"][0].term == "heart attack" + assert response.expansions["synonym"][0].status == GroundingStatus.GROUNDED + assert response.expansions["medication"][0].term == "aspirin" + assert response.expansions["medication"][0].evidence[0].source_id == "C002" + assert response.expansions["symptom"][0].status == GroundingStatus.LLM_ONLY + + +def test_query_expansion_service_can_exclude_llm_only_candidates(tmp_path): + source_file = tmp_path / "terms.yaml" + source_file.write_text("terms: []") + request = QueryExpansionRequest( + term="x", + categories=["symptom"], + llm=LLMConfig(model="test-model"), + grounding=GroundingOptions(include_llm_only=False), + ) + + response = QueryExpansionService(generator=FakeGenerator()).expand( + request, sources=[LocalTerminologySource("local", source_file)] + ) + + assert response.expansions == {"symptom": []} diff --git a/test/rag/test_marqo_rag_utils.py b/test/rag/test_marqo_rag_utils.py new file mode 100644 index 0000000..901a0f8 --- /dev/null +++ b/test/rag/test_marqo_rag_utils.py @@ -0,0 +1,43 @@ +from src.rag.marqo_rag_utils import extract_text_from_highlights + + +def test_extract_text_from_highlights_adds_snippet_offsets_and_metadata(): + highlights, texts, metadata = extract_text_from_highlights( + [ + { + "text": "Alpha beta inflammation gamma.", + "doc_id": "doc-1", + "doc_name": "doc.txt", + "chunk_index": 2, + "chunk_start": 100, + "chunk_end": 130, + "_highlights": [{"text": "inflammation"}], + "_score": 0.9, + } + ], + truncate=False, + ) + + assert highlights == ["inflammation"] + assert texts == ["Alpha beta inflammation gamma."] + assert metadata == [ + { + "doc_id": "doc-1", + "doc_name": "doc.txt", + "chunk_index": 2, + "chunk_start": 100, + "chunk_end": 130, + "retrieved_snippet": "Alpha beta inflammation gamma.", + "retrieved_snippet_start": 0, + "retrieved_snippet_end": 30, + "highlight": "inflammation", + "highlight_field": "text", + "highlight_start": 11, + "highlight_end": 23, + "offset_unit": "retrieved_snippet_char", + "document_highlight_start": 111, + "document_highlight_end": 123, + "document_offset_unit": "document_char", + "retrieved_snippet_index": 0, + } + ] diff --git a/test/rag/test_text_splitters.py b/test/rag/test_text_splitters.py new file mode 100644 index 0000000..f60a58e --- /dev/null +++ b/test/rag/test_text_splitters.py @@ -0,0 +1,33 @@ +from spacy.tokens import Doc +from spacy.vocab import Vocab + +from src.common.spacy_extensions import set_spacy_extensions +from src.rag.text_splitters import PreprocessedSpacyTextSplitter + + +def _doc(text, doc_id="doc-1", offset=0): + words = text.split(" ") + spaces = [True] * (len(words) - 1) + [False] + doc = Doc(Vocab(), words=words, spaces=spaces) + doc._.doc_id = doc_id + doc._.doc_name = "doc.txt" + doc._.offset_in_doc = offset + return doc + + +def test_split_preprocessed_sentences_propagates_document_offsets(): + set_spacy_extensions() + splitter = PreprocessedSpacyTextSplitter(chunk_size=100, chunk_overlap=0) + + chunks, metadata = next( + splitter.split_preprocessed_sentences( + [_doc("Alpha beta", offset=10), _doc("Gamma", offset=30)], + "doc_id", + keep_metadata=["doc_id", "doc_name"], + ) + ) + + assert chunks == ["Alpha beta\n\nGamma"] + assert metadata["doc_id"] == "doc-1" + assert metadata["doc_name"] == "doc.txt" + assert metadata["chunk_offsets"] == [(10, 35)] diff --git a/uv.lock b/uv.lock index bdbe857..ff33201 100644 --- a/uv.lock +++ b/uv.lock @@ -227,7 +227,7 @@ wheels = [ [[package]] name = "concept-graphs" -version = "1.0.0" +version = "1.0.1" source = { virtual = "." } dependencies = [ { name = "altair" }, @@ -263,6 +263,19 @@ dependencies = [ ] [package.dev-dependencies] +llm = [ + { name = "langchain" }, + { name = "langchain-core" }, + { name = "langchain-ollama" }, + { name = "langchain-openai" }, +] +query-expansion = [ + { name = "langchain" }, + { name = "langchain-core" }, + { name = "langchain-ollama" }, + { name = "langchain-openai" }, + { name = "pydantic" }, +] rag = [ { name = "langchain" }, { name = "langchain-community" }, @@ -313,6 +326,19 @@ requires-dist = [ ] [package.metadata.requires-dev] +llm = [ + { name = "langchain" }, + { name = "langchain-core" }, + { name = "langchain-ollama" }, + { name = "langchain-openai" }, +] +query-expansion = [ + { name = "langchain" }, + { name = "langchain-core" }, + { name = "langchain-ollama" }, + { name = "langchain-openai" }, + { name = "pydantic", specifier = ">=2,<3" }, +] rag = [ { name = "langchain" }, { name = "langchain-community" },