From 8b90ccf77956e69c5d0d13b2a33fe9aee21c2745 Mon Sep 17 00:00:00 2001 From: krisnaparahita Date: Fri, 4 Sep 2026 15:10:30 +0800 Subject: [PATCH 1/2] feat: allow setting query transformer in BaseRAGQuestionAnswerer Adds a query_transformer_prompt param that rewrites the query via the LLM (e.g. prompts.prompt_query_rewrite / prompt_query_rewrite_hyde) before retrieval. The rewritten query is used only for document retrieval; the final answer is still generated from the original user prompt, per maintainer feedback on a prior attempt (pathwaycom/pathway#209). Fixes pathwaycom/pathway#67 --- .../pathway/xpacks/llm/question_answering.py | 45 ++++++- python/pathway/xpacks/llm/tests/test_rag.py | 125 ++++++++++++++++++ 2 files changed, 169 insertions(+), 1 deletion(-) diff --git a/python/pathway/xpacks/llm/question_answering.py b/python/pathway/xpacks/llm/question_answering.py index 92be39a2a..239bc1ccb 100644 --- a/python/pathway/xpacks/llm/question_answering.py +++ b/python/pathway/xpacks/llm/question_answering.py @@ -129,6 +129,22 @@ def _get_RAG_prompt_udf(prompt_template: str | Callable[[str, str], str] | pw.UD return verified_template.as_udf() +def _get_query_transformer_prompt_udf( + query_transformer_prompt: Callable[[str], str] | pw.UDF | None, +) -> pw.UDF | None: + if query_transformer_prompt is None: + return None + elif isinstance(query_transformer_prompt, pw.UDF): + return query_transformer_prompt + elif callable(query_transformer_prompt): + return pw.udf(query_transformer_prompt) + else: + raise ValueError( + "Query transformer prompt must be type of one of the following: " + "Callable[[str], str] | ~pw.UDF | None" + ) + + def _get_context_processor_udf( context_processor: ( BaseContextProcessor | Callable[[list[dict] | list[Doc]], str] | pw.UDF @@ -463,6 +479,13 @@ class BaseRAGQuestionAnswerer(SummaryQuestionAnswerer): search_topk: Top k parameter for the retrieval. Adjusts number of chunks in the context. rerank_topk: Number of top-scoring documents to retain after reranking, when a reranker is provided. If ``None``, reranking is disabled. Defaults to ``None``. + query_transformer_prompt: Prompt for transforming the user query before retrieval. + Must be a callable or ``pw.UDF`` that accepts the user query and returns a prompt + to send to ``llm`` for query rewriting. The transformed query is only used for + document retrieval; the answer is still generated from the original user query. + See ``pathway.xpacks.llm.prompts.prompt_query_rewrite`` and + ``pathway.xpacks.llm.prompts.prompt_query_rewrite_hyde`` for ready-to-use options. + Defaults to ``None``, which skips query transformation. Example: @@ -524,6 +547,7 @@ def __init__( search_topk: int = 6, reranker: pw.UDF | None = None, rerank_topk: int | None = None, + query_transformer_prompt: Callable[[str], str] | pw.UDF | None = None, ) -> None: self.llm = llm @@ -536,6 +560,9 @@ def __init__( self._init_schemas(default_llm_name) self.prompt_udf = _get_RAG_prompt_udf(prompt_template) + self.query_transformer_prompt_udf = _get_query_transformer_prompt_udf( + query_transformer_prompt + ) if isinstance(context_processor, BaseContextProcessor): self.docs_to_context_transformer = context_processor.as_udf() @@ -639,11 +666,27 @@ def add_score_to_doc(doc: pw.Json, score: float) -> dict: def answer_query(self, pw_ai_queries: pw.Table) -> pw.Table: """Answer a question based on the available information.""" + if self.query_transformer_prompt_udf is not None: + pw_ai_queries += pw_ai_queries.select( + query_transformer_prompt=self.query_transformer_prompt_udf( + pw.this.prompt + ) + ) + pw_ai_queries += pw_ai_queries.select( + search_query=self.llm( + llms.prompt_chat_single_qa(pw.this.query_transformer_prompt), + model=pw.this.model, + ) + ) + pw_ai_queries = pw_ai_queries.await_futures() + else: + pw_ai_queries += pw_ai_queries.select(search_query=pw.this.prompt) + pw_ai_results = pw_ai_queries + self.indexer.retrieve_query( pw_ai_queries.select( metadata_filter=pw.this.filters, filepath_globpattern=pw.cast(str | None, None), - query=pw.this.prompt, + query=pw.this.search_query, k=self.search_topk, ) ).select( diff --git a/python/pathway/xpacks/llm/tests/test_rag.py b/python/pathway/xpacks/llm/tests/test_rag.py index b93206a0c..17969205c 100644 --- a/python/pathway/xpacks/llm/tests/test_rag.py +++ b/python/pathway/xpacks/llm/tests/test_rag.py @@ -15,6 +15,19 @@ from .utils import build_vector_store, create_rag_app +class _QueryRewriteMockChat(llms.BaseChat): + """Rewrites the rewrite-prompt to a fixed search query; passes other prompts through.""" + + def _accepts_call_arg(self, arg_name: str) -> bool: + return False + + async def __wrapped__(self, messages: list[dict] | pw.Json, model: str) -> str: + content = messages[0]["content"].as_str() + if content.startswith("rewrite:"): + return content.removeprefix("rewrite:") + return model + "," + content + + @pw.udf def fake_embeddings_model(x: str) -> list[float]: return [ @@ -99,6 +112,118 @@ def test_base_rag(): ) +def test_rag_app_set_query_transformer_prompt(): + def query_transformer_prompt(query: str) -> str: + return f"rewrite:{query}" + + rag_app = create_rag_app(query_transformer_prompt=query_transformer_prompt) + + assert isinstance(rag_app.query_transformer_prompt_udf, pw.UDF) + + assert _unwrap_udf(rag_app.query_transformer_prompt_udf)("foo") == "rewrite:foo" + + +def test_rag_app_no_query_transformer_by_default(): + rag_app = create_rag_app() + + assert rag_app.query_transformer_prompt_udf is None + + +def test_base_rag_uses_original_prompt_for_answer_without_transformer(): + schema = pw.schema_from_types(data=bytes, _metadata=dict) + input = pw.debug.table_from_rows( + schema=schema, rows=[("foo", {}), ("bar", {}), ("baz", {})] + ) + + vector_server = VectorStoreServer( + input, + embedder=fake_embeddings_model, + ) + + rag = BaseRAGQuestionAnswerer( + IdentityMockChat(), + vector_server, + prompt_template=_prompt_template, + summarize_template=_summarize_template, + search_topk=1, + ) + + answer_queries = pw.debug.table_from_rows( + schema=rag.AnswerQuerySchema, + rows=[ + ("foo", None, "gpt3.5", False), + ], + ) + + answer_output = rag.answer_query(answer_queries) + + casted_table = answer_output.select( + result=pw.apply_with_type(lambda x: x.value, str, pw.this.result["response"]) + ) + + assert_table_equality( + casted_table, + pw.debug.table_from_markdown( + """ + result + gpt3.5,foo + """ + ), + ) + + +def test_base_rag_query_transformer_used_only_for_retrieval(): + @pw.udf + def query_transformer_prompt(query: str) -> str: + return "rewrite:bar" + + schema = pw.schema_from_types(data=bytes, _metadata=dict) + input = pw.debug.table_from_rows( + schema=schema, rows=[("foo", {}), ("bar", {}), ("baz", {})] + ) + + vector_server = VectorStoreServer( + input, + embedder=fake_embeddings_model, + ) + + rag = BaseRAGQuestionAnswerer( + _QueryRewriteMockChat(), + vector_server, + prompt_template=_prompt_template, + query_transformer_prompt=query_transformer_prompt, + summarize_template=_summarize_template, + search_topk=1, + ) + + answer_queries = pw.debug.table_from_rows( + schema=rag.AnswerQuerySchema, + rows=[ + ("foo", None, "gpt3.5", False), + ], + ) + + answer_output = rag.answer_query(answer_queries) + + casted_table = answer_output.select( + search_query=pw.this.search_query, + result=pw.apply_with_type(lambda x: x.value, str, pw.this.result["response"]), + ) + + # `search_query` shows retrieval used the rewritten query ("bar", not the + # original prompt "foo"), while `result` shows the final LLM call still + # ran on the original, untransformed prompt. + assert_table_equality( + casted_table, + pw.debug.table_from_markdown( + """ + search_query | result + bar | gpt3.5,foo + """ + ), + ) + + def test_rag_app_set_prompt(): prompt_template = "Answer the question. Context: {context}\nQuestion: {query}" From 870712510cb5398be98fe1875e2ad7d5df92f00c Mon Sep 17 00:00:00 2001 From: krisnaparahita Date: Tue, 6 Oct 2026 17:11:10 +0800 Subject: [PATCH 2/2] fix: fall back to original query when RAG rewrite returns None --- .../pathway/xpacks/llm/question_answering.py | 3 +++ python/pathway/xpacks/llm/tests/test_rag.py | 22 +++++++++++-------- 2 files changed, 16 insertions(+), 9 deletions(-) diff --git a/python/pathway/xpacks/llm/question_answering.py b/python/pathway/xpacks/llm/question_answering.py index 239bc1ccb..20a23795a 100644 --- a/python/pathway/xpacks/llm/question_answering.py +++ b/python/pathway/xpacks/llm/question_answering.py @@ -679,6 +679,9 @@ def answer_query(self, pw_ai_queries: pw.Table) -> pw.Table: ) ) pw_ai_queries = pw_ai_queries.await_futures() + pw_ai_queries = pw_ai_queries.with_columns( + search_query=pw.coalesce(pw.this.search_query, pw.this.prompt) + ) else: pw_ai_queries += pw_ai_queries.select(search_query=pw.this.prompt) diff --git a/python/pathway/xpacks/llm/tests/test_rag.py b/python/pathway/xpacks/llm/tests/test_rag.py index 17969205c..585fe38bc 100644 --- a/python/pathway/xpacks/llm/tests/test_rag.py +++ b/python/pathway/xpacks/llm/tests/test_rag.py @@ -21,10 +21,12 @@ class _QueryRewriteMockChat(llms.BaseChat): def _accepts_call_arg(self, arg_name: str) -> bool: return False - async def __wrapped__(self, messages: list[dict] | pw.Json, model: str) -> str: + async def __wrapped__( + self, messages: list[dict] | pw.Json, model: str + ) -> str | None: content = messages[0]["content"].as_str() if content.startswith("rewrite:"): - return content.removeprefix("rewrite:") + return content.removeprefix("rewrite:") or None return model + "," + content @@ -172,10 +174,13 @@ def test_base_rag_uses_original_prompt_for_answer_without_transformer(): ) -def test_base_rag_query_transformer_used_only_for_retrieval(): +@pytest.mark.parametrize("rewritten_query", ["bar", None]) +def test_base_rag_query_transformer_used_only_for_retrieval( + rewritten_query: str | None, +): @pw.udf def query_transformer_prompt(query: str) -> str: - return "rewrite:bar" + return f"rewrite:{rewritten_query or ''}" schema = pw.schema_from_types(data=bytes, _metadata=dict) input = pw.debug.table_from_rows( @@ -210,15 +215,14 @@ def query_transformer_prompt(query: str) -> str: result=pw.apply_with_type(lambda x: x.value, str, pw.this.result["response"]), ) - # `search_query` shows retrieval used the rewritten query ("bar", not the - # original prompt "foo"), while `result` shows the final LLM call still - # ran on the original, untransformed prompt. + # Retrieval uses the rewritten query, falling back to the original prompt + # when the LLM returns None. The final LLM call still uses the original prompt. assert_table_equality( casted_table, pw.debug.table_from_markdown( - """ + f""" search_query | result - bar | gpt3.5,foo + {rewritten_query or 'foo'} | gpt3.5,foo """ ), )