PERI-AI-dev / src /query_engine.py
Nelly-43's picture
Update src/query_engine.py
8a24416 verified
Raw
History Blame Contribute Delete
12.4 kB
"""Query understanding layer: reformulates user queries for better retrieval."""
import json
import os
import re
from typing import Dict
from src.llm import generate
from src.prompts import (build_query_understanding_prompt,
build_categorizer_prompt, build_init_prompt, build_vc_analyzer_prompt)
def analyze_committment_to_vc(value_chain:str,
cfg:Dict,):
domain = cfg.get("chatbot", {}).get("domain", "research")
prompt = build_vc_analyzer_prompt(value_chain, domain)
try:
raw = generate(
"",
prompt,
cfg,
max_tokens=512,
)
return _parse_vc_analyzer_result(raw)
except Exception:
# LLM failure: fall back to raw query
return {
value_chain:{
"financial_commitment": {
"budget_start_year": "",
"budget_start_value": "",
"budget_end_year": "",
"budget_end_value": "",
"currency": "",
"responsible_institution": "",
"evidence": [
{
"quote": "",
"source_document": "",
"page_or_section": ""
}
]
},
"strategic_commitment": {
"documents_examined": 0,
"documents_discussing_value_chain": 0,
"coverage_weight": 0,
"development_plan": {
"value": "NOT_FOUND",
"evidence": []
},
"agricultural_strategy": {
"value": "NOT_FOUND",
"evidence": []
},
"standalone_vc_strategy": {
"value": "NOT_FOUND",
"evidence": []
},
"flagship_project": {
"value": "NOT_FOUND",
"evidence": []
}
},
"institutional_commitment": {
"dedicated_agency": {
"value": "NOT_FOUND",
"evidence": []
},
"dedicated_research_institute": {
"value": "NOT_FOUND",
"evidence": []
},
"dedicated_ministerial_department": {
"value": "NOT_FOUND",
"evidence": []
}
}
}
}
def _load_sql_schema_summary(sql_db_dir: str) -> str:
"""Load SQL schema summary from sql_schemas.json if it exists."""
schema_path = os.path.join(sql_db_dir, "sql_schemas.json")
if not os.path.exists(schema_path):
return ""
try:
with open(schema_path, "r", encoding="utf-8") as f:
schema = json.load(f)
from src.sql_retriever import build_schema_summary
return build_schema_summary(schema)
except Exception:
return ""
def categorize_query(user_query:str,
cfg:Dict,
conversation_history: list[dict] | None = None,):
domain = cfg.get("chatbot", {}).get("domain", "research")
cat_cfg = cfg.get("query_categorization", {})
max_history = cat_cfg.get("max_history", 6)
# Trim conversation history
history = conversation_history or []
if len(history) > max_history:
history = history[-max_history:]
prompt = build_categorizer_prompt(user_query, domain)
try:
raw = generate(
"You are a query-routing classifier. Return only JSON.",
prompt,
cfg,
max_tokens=512,
)
return _parse_qcat_result(raw)
except Exception:
# LLM failure: fall back to raw query
return {
"category": "complete",
"sub-action": "unresolved",
}
def init_query(user_query:str,
cfg:Dict,
conversation_history: list[dict] | None = None,):
domain = cfg.get("chatbot", {}).get("domain", "research")
cat_cfg = cfg.get("query_categorization", {})
max_history = cat_cfg.get("max_history", 6)
# Trim conversation history
history = conversation_history or []
if len(history) > max_history:
history = history[-max_history:]
prompt = build_init_prompt(user_query, domain)
try:
raw = generate(
"You are a query-routing classifier. Return only JSON.",
prompt,
cfg,
max_tokens=512,
)
return _parse_qinit_result(raw)
except Exception:
# LLM failure: fall back to raw query
return {
"country": None,
"value_chain": None,
"investment": None,
}
def understand_query(
user_query: str,
cfg: Dict
) -> dict:
"""Analyze and reformulate a user query for better retrieval.
Uses the LLM to either:
- SEARCH: Rewrite the query into a search-optimized form
- CLARIFY: Ask the user a clarification question
Args:
user_query: Raw user input.
cfg: App config (provides domain, LLM settings).
conversation_history: Recent messages for pronoun/reference resolution.
Returns:
Dict with keys:
- ``action``: ``"search"`` or ``"clarify"``
- ``search_query``: Keyword-optimized query for retrieval (always present)
- ``display_query``: Clear natural-language question for response generation
- ``original_query``: The raw user input
- ``clarification_question``: Question to ask (only if action is "clarify")
"""
domain = cfg.get("chatbot", {}).get("domain", "research")
# Load SQL schema summary for prompt injection
sql_enabled = cfg.get("sql", {}).get("enabled", True)
schema_summary = ""
if sql_enabled:
sql_db_dir = cfg.get("paths", {}).get("sql_db", "sql_db")
if not os.path.isabs(sql_db_dir):
from pathlib import Path as _Path
project_root = _Path(__file__).resolve().parent.parent
sql_db_dir = os.path.join(str(project_root), sql_db_dir)
schema_summary = _load_sql_schema_summary(sql_db_dir)
# Load KB meta overview for routing awareness
from src.kb_meta import load_kb_meta
kb_overview = load_kb_meta(cfg)
prompt = build_query_understanding_prompt(
user_query,
domain,
sql_schema_summary=schema_summary,
kb_overview=kb_overview,
)
try:
raw = generate(
"You are a query reformulation assistant. Return only JSON.",
prompt,
cfg,
max_tokens=2048,
)
result = _parse_qu_result(raw, user_query)
except Exception as e:
# LLM failure: fall back to raw query
result = {
"search_query": user_query,
"sql_query": None,
}
return result
def _parse_vc_analyzer_result(raw:str) -> Dict:
parsed = None
# Try direct parse
try:
parsed = json.loads(raw)
except (json.JSONDecodeError, TypeError):
pass
# Fallback: scan for the first valid JSON object using raw_decode
if parsed is None:
decoder = json.JSONDecoder()
for i, ch in enumerate(raw):
if ch == '{':
try:
parsed, _ = decoder.raw_decode(raw, i)
break
except json.JSONDecodeError:
continue
# Unparseable: fall back to raw query
if not isinstance(parsed, dict):
return {
"value_chain": None,
"financial_commitment": {
"budget_start_year": "",
"budget_start_value": "",
"budget_end_year": "",
"budget_end_value": "",
"currency": "",
"responsible_institution": "",
"evidence": [
{
"quote": "",
"source_document": "",
"page_or_section": ""
}
]
},
"strategic_commitment": {
"documents_examined": 0,
"documents_discussing_value_chain": 0,
"coverage_weight": 0,
"development_plan": {
"value": "NOT_FOUND",
"evidence": []
},
"agricultural_strategy": {
"value": "NOT_FOUND",
"evidence": []
},
"standalone_vc_strategy": {
"value": "NOT_FOUND",
"evidence": []
},
"flagship_project": {
"value": "NOT_FOUND",
"evidence": []
}
},
"institutional_commitment": {
"dedicated_agency": {
"value": "NOT_FOUND",
"evidence": []
},
"dedicated_research_institute": {
"value": "NOT_FOUND",
"evidence": []
},
"dedicated_ministerial_department": {
"value": "NOT_FOUND",
"evidence": []
}
}
}
return parsed
def _parse_qinit_result(raw:str) -> Dict:
parsed = None
# Try direct parse
try:
parsed = json.loads(raw)
except (json.JSONDecodeError, TypeError):
pass
# Fallback: scan for the first valid JSON object using raw_decode
if parsed is None:
decoder = json.JSONDecoder()
for i, ch in enumerate(raw):
if ch == '{':
try:
parsed, _ = decoder.raw_decode(raw, i)
break
except json.JSONDecodeError:
continue
# Unparseable: fall back to raw query
if not isinstance(parsed, dict) or "country" not in parsed:
return {
"country": None,
"value_chain": None,
"investment": None,
}
return parsed
def _parse_qcat_result(raw:str) -> Dict:
parsed = None
# Try direct parse
try:
parsed = json.loads(raw)
except (json.JSONDecodeError, TypeError):
pass
# Fallback: scan for the first valid JSON object using raw_decode
if parsed is None:
decoder = json.JSONDecoder()
for i, ch in enumerate(raw):
if ch == '{':
try:
parsed, _ = decoder.raw_decode(raw, i)
break
except json.JSONDecodeError:
continue
# Unparseable: fall back to raw query
if not isinstance(parsed, dict) or "category" not in parsed:
return {
"category": "general",
"action": "continue",
}
return parsed
def _parse_qu_result(raw: str, original_query: str) -> Dict:
"""Parse the LLM's JSON response into a structured result.
Falls back to raw query passthrough if parsing fails.
"""
parsed = None
# Try direct parse
try:
parsed = json.loads(raw)
except (json.JSONDecodeError, TypeError):
pass
# Fallback: scan for the first valid JSON object using raw_decode
if parsed is None:
decoder = json.JSONDecoder()
for i, ch in enumerate(raw):
if ch == '{':
try:
parsed, _ = decoder.raw_decode(raw, i)
break
except json.JSONDecodeError:
continue
# Unparseable: fall back to raw query
if not isinstance(parsed, dict):
return {
"search_query": original_query,
"display_query": original_query,
"sql_query": None,
}
search_query = parsed.get("search_query", original_query)
if not isinstance(search_query, str):
search_query = original_query
display_query = parsed.get("display_query", original_query)
if not isinstance(display_query, str):
display_query = original_query
sql_query = parsed.get("sql_query") or None
if sql_query is not None and not isinstance(sql_query, str):
sql_query = None
return {
"search_query": search_query,
"display_query": display_query,
"sql_query": sql_query,
}