File size: 7,928 Bytes
dff6d0e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
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()