Rewrite the decontextualized query into simpler keyword queries. The results are stored in handler.rewritten_queries.
(self)
| 28 | self.handler.state.start_precheck_step(self.STEP_NAME) |
| 29 | |
| 30 | async def do(self): |
| 31 | """ |
| 32 | Rewrite the decontextualized query into simpler keyword queries. |
| 33 | The results are stored in handler.rewritten_queries. |
| 34 | """ |
| 35 | # Skip query rewrite for sites that don't support standard retrieval |
| 36 | if not site_supports_standard_retrieval(self.handler.site): |
| 37 | self.handler.rewritten_queries = [self.handler.query] |
| 38 | await self.handler.state.precheck_step_done(self.STEP_NAME) |
| 39 | return |
| 40 | |
| 41 | # Wait for decontextualization to complete since we need the decontextualized query |
| 42 | await self.handler.state._decon_event.wait() |
| 43 | |
| 44 | |
| 45 | try: |
| 46 | # Run the query rewrite prompt |
| 47 | response = await self.run_prompt(self.QUERY_REWRITE_PROMPT_NAME, level="high") |
| 48 | |
| 49 | if not response: |
| 50 | print("No response from QueryRewrite prompt, using original query") |
| 51 | self.handler.rewritten_queries = [self.handler.decontextualized_query] |
| 52 | await self.handler.state.precheck_step_done(self.STEP_NAME) |
| 53 | return |
| 54 | |
| 55 | # Extract the rewritten queries from the response |
| 56 | rewritten_queries = response.get("rewritten_queries", []) |
| 57 | |
| 58 | # Validate the response |
| 59 | if not rewritten_queries or not isinstance(rewritten_queries, list): |
| 60 | print("Invalid response from QueryRewrite prompt, using original query") |
| 61 | self.handler.rewritten_queries = [self.handler.decontextualized_query] |
| 62 | else: |
| 63 | # Filter out any empty queries and ensure they are strings |
| 64 | valid_queries = [q for q in rewritten_queries if q and isinstance(q, str) and q.strip()] |
| 65 | |
| 66 | if not valid_queries: |
| 67 | print("No valid rewritten queries, using original query") |
| 68 | self.handler.rewritten_queries = [self.handler.decontextualized_query] |
| 69 | else: |
| 70 | # Limit to 5 queries maximum |
| 71 | self.handler.rewritten_queries = valid_queries[:5] |
| 72 | print(f"Generated {len(self.handler.rewritten_queries)} rewritten queries: {self.handler.rewritten_queries}") |
| 73 | |
| 74 | # Send a message to the client about the rewritten queries |
| 75 | if hasattr(self.handler, 'rewritten_queries') and len(self.handler.rewritten_queries) > 1: |
| 76 | message = { |
| 77 | "message_type": "query_rewrite", |
| 78 | "original_query": self.handler.decontextualized_query, |
| 79 | "rewritten_queries": self.handler.rewritten_queries, |
| 80 | "query_id": getattr(self.handler, 'query_id', None) |
| 81 | } |
| 82 | asyncio.create_task(self.handler.send_message(message)) |
| 83 | |
| 84 | except Exception as e: |
| 85 | logger.error(f"Error during query rewrite: {e}") |
| 86 | # On error, fall back to using the original query |
| 87 | self.handler.rewritten_queries = [self.handler.decontextualized_query] |
nothing calls this directly
no test coverage detected