RAG2 / pseudo_agent.py
antimoda1
fix agent
dff6d0e
Raw
History Blame Contribute Delete
7.93 kB
import re
from llm import get_llm_answer, get_llm_completion
from retrieval_impl import RETRIEVAL
from vocabulary.parse_vocabulary import VOCABULARY_MANAGER
MAX_ITERATIONS = 3
DEFAULT_YEAR_FROM = 1918
DEFAULT_YEAR_TO = 2026
DEFAULT_TOP_K = 20
EXPANDED_TOP_K = 30
NEED_MORE_RE = re.compile(r"^NEED_MORE:\s*(.+)$", re.MULTILINE)
MULTI_RESULT_HINTS = ("факт", "пример", "несколько", "пять", "четыре", "три", "подробн")
def _wants_more_results(query: str) -> bool:
q = query.lower()
return any(hint in q for hint in MULTI_RESULT_HINTS)
def _merge_indices(existing: list, new: list) -> list:
seen = set(existing)
merged = list(existing)
for idx in new:
if idx not in seen:
merged.append(idx)
seen.add(idx)
return merged
def _format_context(indices: list) -> str:
if not indices:
return ""
df = RETRIEVAL.paragraphs_df.iloc[indices]
chunks = []
for _, row in df.iterrows():
chunks.append(
f"""Название: {row.summary}
Период: {row.start_year}-{row.end_year}
{row.text}"""
)
return "\n\n---\n\n".join(chunks)
def _find_title_boost_query(query: str) -> str | None:
"""Если запрос пересекается с названием раздела — вернуть его для приоритетного поиска."""
q = query.lower()
best_match = None
best_len = 0
for summary in RETRIEVAL.paragraphs_df["summary"].unique():
summary_lower = summary.lower()
if summary_lower in q or q in summary_lower:
if len(summary_lower) > best_len:
best_match = summary
best_len = len(summary_lower)
continue
words = [w for w in re.findall(r"\w+", summary_lower) if len(w) > 3]
if words and all(w in q for w in words[:2]):
if len(summary_lower) > best_len:
best_match = summary
best_len = len(summary_lower)
return best_match
def _search(
query: str,
*,
year_from: int,
year_to: int,
top_k: int,
) -> tuple[list, str]:
_, _, indices, status = RETRIEVAL.perform_search(
query=query,
top_k=top_k,
year_from=year_from,
year_to=year_to,
)
return list(indices), status
def _build_judge_prompt(user_query: str, context: str) -> str:
return f"""Ты оцениваешь, достаточно ли архивных материалов для ответа на вопрос пользователя.
ВОПРОС:
{user_query}
НАЙДЕННЫЕ МАТЕРИАЛЫ:
{context if context else "(пусто — материалов нет)"}
Правила:
- Смотри ТОЛЬКО на материалы, не на свои знания.
- Корпус — только про общественный транспорт Рязани.
- Если материалов хватает для ответа — ответь ровно одной строкой:
SUFFICIENT
- Если материалов мало, они не по теме или их нет — ответь ровно одной строкой:
NEED_MORE: <короткий поисковый запрос на русском>
- Никакого другого текста.
Примеры:
Вопрос: "история маршрутов с буквами в номерах"
Материалы: [текст про буквы в номерах маршрутов]
→ SUFFICIENT
Вопрос: "пять интересных фактов об истории транспорта"
Материалы: [один короткий абзац]
→ NEED_MORE: интересные факты история транспорта Рязань"""
def _parse_judge_response(text: str) -> tuple[str, str | None]:
text = text.strip()
for line in text.splitlines():
line = line.strip()
if line == "SUFFICIENT":
return "sufficient", None
match = NEED_MORE_RE.match(line)
if match:
return "need_more", match.group(1).strip()
if "SUFFICIENT" in text:
return "sufficient", None
match = NEED_MORE_RE.search(text)
if match:
return "need_more", match.group(1).strip()
return "sufficient", None
class PseudoAgent:
def __init__(
self,
*,
year_from: int = DEFAULT_YEAR_FROM,
year_to: int = DEFAULT_YEAR_TO,
max_iterations: int = MAX_ITERATIONS,
):
self.year_from = year_from
self.year_to = year_to
self.max_iterations = max_iterations
def run(self, query: str):
"""Генератор: прогресс поиска/судьи и финальный ответ."""
query = query.strip()
if not query:
yield "Введите вопрос"
return
parts = ["**Псевдо-агент:** поиск материалов...\n"]
yield "\n\n".join(parts)
all_indices: list = []
search_query = query
top_k = EXPANDED_TOP_K if _wants_more_results(query) else DEFAULT_TOP_K
title_boost = _find_title_boost_query(query)
if title_boost:
boost_indices, boost_status = _search(
title_boost,
year_from=self.year_from,
year_to=self.year_to,
top_k=top_k,
)
all_indices = _merge_indices(all_indices, boost_indices)
parts.append(f"**Совпадение с разделом:** `{title_boost}` — {boost_status}")
yield "\n\n".join(parts)
for iteration in range(1, self.max_iterations + 1):
parts.append(f"### Итерация {iteration}")
indices, status = _search(
search_query,
year_from=self.year_from,
year_to=self.year_to,
top_k=top_k,
)
all_indices = _merge_indices(all_indices, indices)
parts.append(f"**Поиск:** `{search_query}` (top_k={top_k})")
parts.append(f"**Результат:** {status}, всего разделов: {len(all_indices)}")
yield "\n\n".join(parts)
context = _format_context(all_indices)
if not context:
if iteration < self.max_iterations:
search_query = query
top_k = EXPANDED_TOP_K
parts.append("**Судья:** (пропущен — материалов нет, повторяю поиск)")
yield "\n\n".join(parts)
continue
break
judge_prompt = _build_judge_prompt(query, context)
judge_raw = get_llm_completion(judge_prompt, max_tokens=100, temperature=0.0)
decision, next_query = _parse_judge_response(judge_raw)
parts.append(f"**Судья:** `{judge_raw.strip()}`")
yield "\n\n".join(parts)
if decision == "sufficient" or iteration == self.max_iterations:
break
if next_query:
search_query = next_query
else:
search_query = query
context = _format_context(all_indices)
if not context:
yield "\n\n".join(parts) + "\n\n## Ответ\n\nВ корпусе не нашёл материалов по этому запросу. Попробуйте уточнить вопрос."
return
parts.append("## Ответ\n\n")
yield "\n\n".join(parts)
answer_prompt = VOCABULARY_MANAGER.wrap_prompt(context, query)
answer_text = ""
for chunk in get_llm_answer(answer_prompt):
answer_text += chunk
yield "\n\n".join(parts[:-1]) + parts[-1] + answer_text
PSEUDO_AGENT = PseudoAgent()