Files
voice-call-center-ai/app/knowledge.py
T
2026-06-20 17:38:19 +05:00

268 lines
7.6 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
from __future__ import annotations
import json
import math
import os
import re
from collections import Counter
from dataclasses import dataclass
from pathlib import Path
from typing import Any
TOKEN_RE = re.compile(r"[0-9A-Za-zА-Яа-яЁёӘәҒғҚқҢңӨөҰұҮүҺһІі]+", re.UNICODE)
WHITESPACE_RE = re.compile(r"\s+")
STOPWORDS = {
"а",
"без",
"бы",
"в",
"во",
"где",
"для",
"до",
"его",
"ее",
"если",
"за",
"и",
"или",
"как",
"какая",
"какие",
"какой",
"когда",
"кто",
"ли",
"мне",
"можно",
"мы",
"на",
"надо",
"не",
"нет",
"но",
"о",
"об",
"по",
"подскажите",
"пожалуйста",
"при",
"с",
"со",
"то",
"у",
"что",
"чтобы",
"это",
"я",
}
@dataclass(frozen=True)
class KnowledgeHit:
score: float
record: dict[str, Any]
@dataclass
class KnowledgeDocument:
record: dict[str, Any]
title_tokens: set[str]
keyword_tokens: set[str]
content_tokens: Counter[str]
normalized_text: str
class KnowledgeBase:
def __init__(self, documents: list[KnowledgeDocument], source_files: list[str]):
self.documents = documents
self.source_files = source_files
self._idf = self._build_idf(documents)
@classmethod
def empty(cls) -> "KnowledgeBase":
return cls([], [])
@classmethod
def from_directory(cls, directory: str) -> "KnowledgeBase":
root = Path(directory)
if not root.exists():
return cls.empty()
records: list[dict[str, Any]] = []
source_files: list[str] = []
for path in sorted(root.glob("*.jsonl")):
source_files.append(str(path))
with path.open("r", encoding="utf-8") as file:
for line in file:
line = line.strip()
if not line:
continue
records.append(json.loads(line))
documents = [build_document(record) for record in records]
return cls(documents, source_files)
@property
def size(self) -> int:
return len(self.documents)
def search(self, query: str, limit: int = 4, min_score: float = 2.0) -> list[KnowledgeHit]:
query_tokens = tokenize(query)
if not query_tokens:
return []
query_text = normalize_text(query)
hits: list[KnowledgeHit] = []
for document in self.documents:
score = self._score_document(document, query_tokens, query_text)
if score >= min_score:
hits.append(KnowledgeHit(score=score, record=document.record))
hits.sort(key=lambda hit: hit.score, reverse=True)
return hits[:limit]
def _score_document(
self,
document: KnowledgeDocument,
query_tokens: list[str],
query_text: str,
) -> float:
unique_query_tokens = set(query_tokens)
score = 0.0
for token in unique_query_tokens:
idf = self._idf.get(token, 1.0)
if token in document.title_tokens:
score += 3.4 * idf
if token in document.keyword_tokens:
score += 2.6 * idf
count = document.content_tokens.get(token, 0)
if count:
score += min(count, 3) * idf
if query_text and query_text in document.normalized_text:
score += 8.0
elif len(unique_query_tokens) >= 3:
overlap = len(unique_query_tokens & set(document.content_tokens))
score += overlap / len(unique_query_tokens)
return score
@staticmethod
def _build_idf(documents: list[KnowledgeDocument]) -> dict[str, float]:
document_count = max(len(documents), 1)
document_frequency: Counter[str] = Counter()
for document in documents:
document_frequency.update(set(document.content_tokens))
return {
token: math.log((document_count + 1) / (frequency + 1)) + 1
for token, frequency in document_frequency.items()
}
def build_document(record: dict[str, Any]) -> KnowledgeDocument:
title_text = " ".join(
filter(
None,
[
record.get("title", ""),
record.get("question", ""),
record.get("category", ""),
record.get("region_code", ""),
],
)
)
keywords = record.get("keywords") or []
keyword_text = " ".join(str(keyword) for keyword in keywords)
content_text = " ".join(
filter(
None,
[
record.get("content", ""),
record.get("short_answer", ""),
record.get("full_answer", ""),
keyword_text,
],
)
)
return KnowledgeDocument(
record=record,
title_tokens=set(tokenize(title_text)),
keyword_tokens=set(tokenize(keyword_text)),
content_tokens=Counter(tokenize(content_text)),
normalized_text=normalize_text(f"{title_text} {content_text}"),
)
def format_knowledge_context(hits: list[KnowledgeHit], max_chars: int = 4200) -> str:
if not hits:
return ""
parts = [
"Служебный контекст базы знаний Aimaq. Используй только релевантные факты из этого блока. "
"Не называй клиенту ID записей и внутренние источники."
]
current_len = len(parts[0])
for index, hit in enumerate(hits, start=1):
record = hit.record
entry = format_record(index, hit.score, record)
if current_len + len(entry) > max_chars:
break
parts.append(entry)
current_len += len(entry)
return "\n\n".join(parts)
def format_record(index: int, score: float, record: dict[str, Any]) -> str:
title = record.get("title") or record.get("question") or "Без названия"
category = record.get("category") or "general"
region = record.get("region_code") or "global"
short_answer = record.get("short_answer") or ""
full_answer = record.get("full_answer") or ""
content = record.get("content") or ""
body = full_answer or short_answer or content
if short_answer and full_answer and short_answer not in full_answer:
body = f"{short_answer}\n{full_answer}"
body = trim_text(body, 1100)
return (
f"[KB {index}] score={score:.2f}\n"
f"ID: {record.get('external_id', '')}\n"
f"Категория: {category}\n"
f"Регион: {region}\n"
f"Заголовок: {title}\n"
f"Ответ/факты: {body}"
)
def tokenize(text: str) -> list[str]:
tokens = []
for token in TOKEN_RE.findall(text.casefold().replace("ё", "е")):
if len(token) < 2 or token in STOPWORDS:
continue
tokens.append(token)
return tokens
def normalize_text(text: str) -> str:
text = text.casefold().replace("ё", "е")
text = " ".join(TOKEN_RE.findall(text))
return WHITESPACE_RE.sub(" ", text).strip()
def trim_text(text: str, max_chars: int) -> str:
text = WHITESPACE_RE.sub(" ", text).strip()
if len(text) <= max_chars:
return text
return text[: max_chars - 1].rstrip() + "…"
def resolve_knowledge_dir(project_dir: str, configured_dir: str) -> str:
if os.path.isabs(configured_dir):
return configured_dir
return os.path.join(project_dir, configured_dir)