forked from zovos/vk_hackathon
687 lines
21 KiB
Python
687 lines
21 KiB
Python
import logging
|
||
import os
|
||
import json
|
||
import re
|
||
from dataclasses import dataclass
|
||
from datetime import datetime, timezone
|
||
from functools import lru_cache
|
||
from typing import Any
|
||
import asyncio
|
||
|
||
from fastapi import FastAPI, Request
|
||
from fastapi.exceptions import RequestValidationError
|
||
from fastapi.responses import JSONResponse
|
||
from pydantic import BaseModel
|
||
|
||
# Ваш сервис должен считывать эти переменные из окружения (env), так как проверяющая система управляет ими
|
||
HOST = os.getenv("HOST", "0.0.0.0")
|
||
PORT = int(os.getenv("PORT", "8004"))
|
||
|
||
logging.basicConfig(level=os.getenv("LOG_LEVEL", "INFO"))
|
||
logger = logging.getLogger("index-service")
|
||
|
||
|
||
# Модель данных, которую мы предоставляем и рассчитываем получать от вас
|
||
class Chat(BaseModel):
|
||
id: str
|
||
name: str
|
||
sn: str
|
||
type: str # group, channel, private
|
||
is_public: bool | None = None
|
||
members_count: int | None = None
|
||
members: list[dict[str, Any]] | None = None
|
||
|
||
|
||
class Message(BaseModel):
|
||
id: str
|
||
thread_sn: str | None = None
|
||
time: int
|
||
text: str
|
||
sender_id: str
|
||
file_snippets: str
|
||
parts: list[dict[str, Any]] | None = None
|
||
mentions: list[str] | None = None
|
||
member_event: dict[str, Any] | None = None
|
||
is_system: bool
|
||
is_hidden: bool
|
||
is_forward: bool
|
||
is_quote: bool
|
||
|
||
|
||
class ChatData(BaseModel):
|
||
chat: Chat
|
||
overlap_messages: list[Message]
|
||
new_messages: list[Message]
|
||
|
||
|
||
class IndexAPIRequest(BaseModel):
|
||
data: ChatData
|
||
|
||
|
||
# dense_content будет передан в dense embedding модель для построения семантического вектора.
|
||
# sparse_content будет передан в sparse модель для построения разреженного индекса "по словам".
|
||
# Можно оставить dense_content и sparse_content равными page_content,
|
||
# а можно формировать для них разные версии текста.
|
||
class IndexAPIItem(BaseModel):
|
||
page_content: str
|
||
dense_content: str
|
||
sparse_content: str
|
||
message_ids: list[str]
|
||
|
||
|
||
class IndexAPIResponse(BaseModel):
|
||
results: list[IndexAPIItem]
|
||
|
||
|
||
class SparseEmbeddingRequest(BaseModel):
|
||
texts: list[str]
|
||
|
||
|
||
class SparseVector(BaseModel):
|
||
indices: list[int]
|
||
values: list[float]
|
||
|
||
|
||
class SparseEmbeddingResponse(BaseModel):
|
||
vectors: list[SparseVector]
|
||
|
||
|
||
app = FastAPI(title="Index Service", version="0.1.0")
|
||
|
||
# Ваша внутренняя логика построения чанков. Можете делать всё, что посчитаете нужным.
|
||
# Текущий код – минимальный пример
|
||
|
||
CHUNK_MAX_CHARS = int(os.getenv("INDEX_CHUNK_MAX_CHARS", "2200"))
|
||
MESSAGE_MAX_CHARS = int(os.getenv("INDEX_MESSAGE_MAX_CHARS", "1400"))
|
||
TEXT_SECTION_MAX_CHARS = int(os.getenv("INDEX_TEXT_SECTION_MAX_CHARS", "1000"))
|
||
TIME_GAP_SECONDS = int(os.getenv("INDEX_TIME_GAP_SECONDS", str(6 * 60 * 60)))
|
||
CHUNK_OVERLAP_MESSAGES = int(os.getenv("INDEX_CHUNK_OVERLAP_MESSAGES", "1"))
|
||
SPARSE_MODEL_NAME = "Qdrant/bm25"
|
||
FASTEMBED_CACHE_PATH = "/models/fastembed"
|
||
|
||
# Важная переманная, которая позволяет вычислять sparse вектор в несколько ядер. Не рекомендуется изменять.
|
||
UVICORN_WORKERS = 8
|
||
|
||
|
||
@dataclass(frozen=True)
|
||
class RenderedText:
|
||
page: str
|
||
dense: str
|
||
sparse: str
|
||
|
||
|
||
@dataclass(frozen=True)
|
||
class RenderedUnit:
|
||
message_id: str
|
||
time: int
|
||
text: RenderedText
|
||
|
||
|
||
@dataclass(frozen=True)
|
||
class ChunkUnit:
|
||
unit: RenderedUnit
|
||
is_new: bool
|
||
|
||
|
||
def clean_text(value: Any) -> str:
|
||
if not isinstance(value, str):
|
||
return ""
|
||
|
||
text = (
|
||
value.replace("\r\n", "\n")
|
||
.replace("\r", "\n")
|
||
.replace("\u200b", " ")
|
||
.replace("\xa0", " ")
|
||
)
|
||
lines = [re.sub(r"[ \t]+", " ", line).strip() for line in text.split("\n")]
|
||
result: list[str] = []
|
||
previous_blank = False
|
||
|
||
for line in lines:
|
||
if not line:
|
||
if result and not previous_blank:
|
||
result.append("")
|
||
previous_blank = True
|
||
continue
|
||
|
||
result.append(line)
|
||
previous_blank = False
|
||
|
||
return "\n".join(result).strip()
|
||
|
||
|
||
def unique_preserve_order(values: list[str]) -> list[str]:
|
||
seen: set[str] = set()
|
||
result: list[str] = []
|
||
|
||
for value in values:
|
||
if value and value not in seen:
|
||
seen.add(value)
|
||
result.append(value)
|
||
|
||
return result
|
||
|
||
|
||
def format_time(timestamp: int) -> str:
|
||
return datetime.fromtimestamp(timestamp, tz=timezone.utc).isoformat().replace("+00:00", "Z")
|
||
|
||
|
||
def format_optional_time(value: Any) -> str:
|
||
if value is None or value == "":
|
||
return ""
|
||
|
||
try:
|
||
return format_time(int(value))
|
||
except (TypeError, ValueError, OSError, OverflowError):
|
||
return clean_text(str(value))
|
||
|
||
|
||
def lexical_terms(value: str) -> str:
|
||
return clean_text(re.sub(r"[^0-9A-Za-zА-Яа-яЁё]+", " ", value))
|
||
|
||
|
||
def split_long_text(text: str, max_chars: int) -> list[str]:
|
||
text = clean_text(text)
|
||
if not text:
|
||
return []
|
||
if len(text) <= max_chars:
|
||
return [text]
|
||
|
||
paragraphs = [item.strip() for item in re.split(r"\n{2,}", text) if item.strip()]
|
||
pieces: list[str] = []
|
||
current = ""
|
||
|
||
def append_current() -> None:
|
||
nonlocal current
|
||
if current:
|
||
pieces.append(current)
|
||
current = ""
|
||
|
||
def split_oversized(paragraph: str) -> list[str]:
|
||
words = paragraph.split()
|
||
result: list[str] = []
|
||
part = ""
|
||
|
||
for word in words:
|
||
if not part:
|
||
part = word
|
||
continue
|
||
|
||
if len(part) + 1 + len(word) <= max_chars:
|
||
part += " " + word
|
||
else:
|
||
result.append(part)
|
||
part = word
|
||
|
||
if part:
|
||
result.append(part)
|
||
|
||
return result
|
||
|
||
for paragraph in paragraphs:
|
||
candidates = [paragraph] if len(paragraph) <= max_chars else split_oversized(paragraph)
|
||
for candidate in candidates:
|
||
separator = "\n\n" if current else ""
|
||
if current and len(current) + len(separator) + len(candidate) > max_chars:
|
||
append_current()
|
||
|
||
current = candidate if not current else current + separator + candidate
|
||
|
||
append_current()
|
||
return pieces
|
||
|
||
|
||
def combine_texts(items: list[RenderedText], separator: str = "\n") -> RenderedText:
|
||
return RenderedText(
|
||
page=separator.join(item.page for item in items if item.page).strip(),
|
||
dense=separator.join(item.dense for item in items if item.dense).strip(),
|
||
sparse=separator.join(item.sparse for item in items if item.sparse).strip(),
|
||
)
|
||
|
||
|
||
def rendered_length(text: RenderedText) -> int:
|
||
return max(len(text.page), len(text.dense), len(text.sparse))
|
||
|
||
|
||
def message_header(message: Message) -> RenderedText:
|
||
timestamp = format_time(message.time)
|
||
mentions = unique_preserve_order(message.mentions or [])
|
||
|
||
page_lines = [
|
||
f"author: {message.sender_id}",
|
||
f"time: {timestamp}",
|
||
]
|
||
dense_lines = [
|
||
f"author: {message.sender_id}",
|
||
f"message_time: {timestamp}",
|
||
]
|
||
sparse_terms = [
|
||
"author",
|
||
message.sender_id,
|
||
lexical_terms(message.sender_id),
|
||
timestamp,
|
||
]
|
||
|
||
if message.thread_sn:
|
||
page_lines.append(f"thread: {message.thread_sn}")
|
||
dense_lines.append(f"thread: {message.thread_sn}")
|
||
sparse_terms.extend(["thread", message.thread_sn, lexical_terms(message.thread_sn)])
|
||
|
||
if mentions:
|
||
mentions_text = ", ".join(mentions)
|
||
page_lines.append(f"mentions: {mentions_text}")
|
||
dense_lines.append(f"mentions: {mentions_text}")
|
||
sparse_terms.extend(["mentions", *mentions, *(lexical_terms(item) for item in mentions)])
|
||
|
||
flags = []
|
||
if message.is_system:
|
||
flags.append("system")
|
||
if message.is_forward:
|
||
flags.append("forward")
|
||
if message.is_quote:
|
||
flags.append("quote")
|
||
if message.is_hidden:
|
||
flags.append("hidden")
|
||
|
||
if flags:
|
||
flags_text = ", ".join(flags)
|
||
dense_lines.append(f"message_flags: {flags_text}")
|
||
sparse_terms.extend(flags)
|
||
|
||
return RenderedText(
|
||
page="\n".join(page_lines),
|
||
dense="\n".join(dense_lines),
|
||
sparse=" ".join(term for term in sparse_terms if term),
|
||
)
|
||
|
||
|
||
def render_plain_sections(text: str, label: str = "text") -> list[RenderedText]:
|
||
sections: list[RenderedText] = []
|
||
pieces = split_long_text(text, TEXT_SECTION_MAX_CHARS)
|
||
|
||
for index, piece in enumerate(pieces):
|
||
suffix = f" part {index + 1}/{len(pieces)}" if len(pieces) > 1 else ""
|
||
sections.append(
|
||
RenderedText(
|
||
page=piece,
|
||
dense=f"{label}{suffix}: {piece}",
|
||
sparse=f"{label} {piece}",
|
||
)
|
||
)
|
||
|
||
return sections
|
||
|
||
|
||
def render_part_sections(part: dict[str, Any]) -> list[RenderedText]:
|
||
text = clean_text(part.get("text"))
|
||
if not text:
|
||
return []
|
||
|
||
media_type = clean_text(part.get("mediaType") or part.get("type") or "text").lower()
|
||
source = clean_text(part.get("sn"))
|
||
part_time = part.get("time")
|
||
source_bits = []
|
||
if source:
|
||
source_bits.append(f"source: {source}")
|
||
formatted_part_time = format_optional_time(part_time)
|
||
if formatted_part_time:
|
||
source_bits.append(f"source_time: {formatted_part_time}")
|
||
source_text = ", ".join(source_bits)
|
||
|
||
if media_type == "quote":
|
||
label = f"quote from {source}" if source else "quote"
|
||
dense_label = f"quote; {source_text}" if source_text else "quote"
|
||
sparse_prefix = f"quote цитата {source} {lexical_terms(source)}"
|
||
elif media_type == "forward":
|
||
label = f"forwarded from {source}" if source else "forwarded"
|
||
dense_label = f"forwarded_message; {source_text}" if source_text else "forwarded_message"
|
||
sparse_prefix = f"forward forwarded_message пересланное {source} {lexical_terms(source)}"
|
||
else:
|
||
label = "text"
|
||
dense_label = "text"
|
||
sparse_prefix = "text"
|
||
|
||
sections: list[RenderedText] = []
|
||
pieces = split_long_text(text, TEXT_SECTION_MAX_CHARS)
|
||
|
||
for index, piece in enumerate(pieces):
|
||
suffix = f" part {index + 1}/{len(pieces)}" if len(pieces) > 1 else ""
|
||
page_prefix = f"{label}{suffix}:"
|
||
sections.append(
|
||
RenderedText(
|
||
page=f"{page_prefix}\n{piece}" if media_type in {"quote", "forward"} else piece,
|
||
dense=f"{dense_label}{suffix}: {piece}",
|
||
sparse=f"{sparse_prefix} {piece}",
|
||
)
|
||
)
|
||
|
||
return sections
|
||
|
||
|
||
def render_file_sections(raw_snippets: str) -> list[RenderedText]:
|
||
raw_snippets = clean_text(raw_snippets)
|
||
if not raw_snippets:
|
||
return []
|
||
|
||
try:
|
||
parsed = json.loads(raw_snippets)
|
||
except json.JSONDecodeError:
|
||
return [
|
||
RenderedText(
|
||
page=f"file_snippet: {raw_snippets}",
|
||
dense=f"file_snippet: {raw_snippets}",
|
||
sparse=f"file file_snippet {raw_snippets}",
|
||
)
|
||
]
|
||
|
||
snippets = parsed if isinstance(parsed, list) else [parsed]
|
||
sections: list[RenderedText] = []
|
||
|
||
for snippet in snippets:
|
||
if not isinstance(snippet, dict):
|
||
text = clean_text(str(snippet))
|
||
sections.append(RenderedText(page=f"file: {text}", dense=f"file: {text}", sparse=f"file {text}"))
|
||
continue
|
||
|
||
name = clean_text(snippet.get("name"))
|
||
mime = clean_text(snippet.get("mime"))
|
||
url = clean_text(snippet.get("original_url") or snippet.get("url"))
|
||
owner = clean_text(snippet.get("uid"))
|
||
created = clean_text(snippet.get("date_create"))
|
||
file_bits = [
|
||
f"name: {name}" if name else "",
|
||
f"mime: {mime}" if mime else "",
|
||
f"url: {url}" if url else "",
|
||
f"owner: {owner}" if owner else "",
|
||
f"created: {created}" if created else "",
|
||
]
|
||
file_text = ", ".join(bit for bit in file_bits if bit)
|
||
sparse_terms = " ".join(
|
||
term
|
||
for term in [
|
||
"file",
|
||
"attachment",
|
||
"document",
|
||
name,
|
||
lexical_terms(name),
|
||
mime,
|
||
url,
|
||
owner,
|
||
lexical_terms(owner),
|
||
created,
|
||
]
|
||
if term
|
||
)
|
||
sections.append(
|
||
RenderedText(
|
||
page=f"file: {file_text}",
|
||
dense=f"file: {file_text}",
|
||
sparse=sparse_terms,
|
||
)
|
||
)
|
||
|
||
return sections
|
||
|
||
|
||
def render_member_event(message: Message) -> list[RenderedText]:
|
||
event = message.member_event
|
||
if not event:
|
||
return []
|
||
|
||
event_type = clean_text(event.get("type") or "member_event")
|
||
members_raw = event.get("members")
|
||
members = [clean_text(item) for item in members_raw] if isinstance(members_raw, list) else []
|
||
members = unique_preserve_order([item for item in members if item])
|
||
|
||
if members:
|
||
members_text = ", ".join(members)
|
||
else:
|
||
members_text = " ".join(clean_text(str(value)) for value in event.values() if value)
|
||
|
||
page = f"system_event: {event_type}; actor: {message.sender_id}; members: {members_text}"
|
||
dense = (
|
||
f"system_event: {event_type}; action: add or update chat members; "
|
||
f"actor: {message.sender_id}; members: {members_text}"
|
||
)
|
||
sparse = " ".join(
|
||
term
|
||
for term in [
|
||
"system_event",
|
||
"member_event",
|
||
event_type,
|
||
"addMembers",
|
||
"добавление участников",
|
||
message.sender_id,
|
||
lexical_terms(message.sender_id),
|
||
members_text,
|
||
lexical_terms(members_text),
|
||
]
|
||
if term
|
||
)
|
||
|
||
return [RenderedText(page=page, dense=dense, sparse=sparse)]
|
||
|
||
|
||
def render_message_sections(message: Message) -> list[RenderedText]:
|
||
sections: list[RenderedText] = []
|
||
sections.extend(render_plain_sections(message.text, "message_text"))
|
||
|
||
for part in message.parts or []:
|
||
if isinstance(part, dict):
|
||
sections.extend(render_part_sections(part))
|
||
|
||
sections.extend(render_file_sections(message.file_snippets))
|
||
sections.extend(render_member_event(message))
|
||
|
||
return sections
|
||
|
||
|
||
def render_message_units(message: Message) -> list[RenderedUnit]:
|
||
sections = render_message_sections(message)
|
||
if not sections:
|
||
return []
|
||
|
||
header = message_header(message)
|
||
units: list[RenderedUnit] = []
|
||
current: list[RenderedText] = []
|
||
|
||
def flush() -> None:
|
||
nonlocal current
|
||
if not current:
|
||
return
|
||
|
||
text = combine_texts([header, *current])
|
||
units.append(RenderedUnit(message_id=message.id, time=message.time, text=text))
|
||
current = []
|
||
|
||
for section in sections:
|
||
candidate = combine_texts([header, *current, section])
|
||
if current and rendered_length(candidate) > MESSAGE_MAX_CHARS:
|
||
flush()
|
||
|
||
current.append(section)
|
||
|
||
flush()
|
||
return units
|
||
|
||
|
||
def chunk_text(items: list[ChunkUnit]) -> RenderedText:
|
||
return RenderedText(
|
||
page="\n\n".join(item.unit.text.page for item in items if item.unit.text.page).strip(),
|
||
dense="\n\n".join(item.unit.text.dense for item in items if item.unit.text.dense).strip(),
|
||
sparse="\n\n".join(item.unit.text.sparse for item in items if item.unit.text.sparse).strip(),
|
||
)
|
||
|
||
|
||
def chunk_length(items: list[ChunkUnit]) -> int:
|
||
return rendered_length(chunk_text(items))
|
||
|
||
|
||
def trim_context(context: list[RenderedUnit], unit: RenderedUnit) -> list[RenderedUnit]:
|
||
result = context[-CHUNK_OVERLAP_MESSAGES:] if CHUNK_OVERLAP_MESSAGES > 0 else []
|
||
items = [ChunkUnit(item, False) for item in result] + [ChunkUnit(unit, True)]
|
||
|
||
while result and chunk_length(items) > CHUNK_MAX_CHARS:
|
||
result = result[1:]
|
||
items = [ChunkUnit(item, False) for item in result] + [ChunkUnit(unit, True)]
|
||
|
||
return result
|
||
|
||
|
||
def build_chunks(
|
||
overlap_messages: list[Message],
|
||
new_messages: list[Message],
|
||
) -> list[IndexAPIItem]:
|
||
result: list[IndexAPIItem] = []
|
||
|
||
overlap_units = [
|
||
unit
|
||
for message in overlap_messages
|
||
for unit in render_message_units(message)
|
||
]
|
||
new_units = [
|
||
unit
|
||
for message in new_messages
|
||
for unit in render_message_units(message)
|
||
]
|
||
current: list[ChunkUnit] = []
|
||
last_new_time: int | None = None
|
||
|
||
def flush_current() -> None:
|
||
nonlocal current
|
||
if not current:
|
||
return
|
||
|
||
message_ids = unique_preserve_order(
|
||
[item.unit.message_id for item in current if item.is_new]
|
||
)
|
||
if not message_ids:
|
||
current = []
|
||
return
|
||
|
||
text = chunk_text(current)
|
||
result.append(
|
||
IndexAPIItem(
|
||
page_content=text.page,
|
||
dense_content=text.dense,
|
||
sparse_content=text.sparse,
|
||
message_ids=message_ids,
|
||
)
|
||
)
|
||
current = []
|
||
|
||
def request_overlap_context(unit: RenderedUnit) -> list[RenderedUnit]:
|
||
close_units = [
|
||
item
|
||
for item in overlap_units
|
||
if abs(unit.time - item.time) <= TIME_GAP_SECONDS
|
||
]
|
||
return trim_context(close_units, unit)
|
||
|
||
for unit in new_units:
|
||
if not current:
|
||
context = request_overlap_context(unit)
|
||
current = [ChunkUnit(item, False) for item in context]
|
||
current.append(ChunkUnit(unit, True))
|
||
last_new_time = unit.time
|
||
continue
|
||
|
||
gap = abs(unit.time - last_new_time) if last_new_time is not None else 0
|
||
candidate = [*current, ChunkUnit(unit, True)]
|
||
should_split = gap > TIME_GAP_SECONDS or chunk_length(candidate) > CHUNK_MAX_CHARS
|
||
|
||
if should_split:
|
||
previous_new_units = [item.unit for item in current if item.is_new]
|
||
context = (
|
||
trim_context(previous_new_units, unit)
|
||
if gap <= TIME_GAP_SECONDS
|
||
else request_overlap_context(unit)
|
||
)
|
||
flush_current()
|
||
current = [ChunkUnit(item, False) for item in context]
|
||
current.append(ChunkUnit(unit, True))
|
||
else:
|
||
current.append(ChunkUnit(unit, True))
|
||
|
||
last_new_time = unit.time
|
||
|
||
flush_current()
|
||
|
||
return result
|
||
|
||
# Ваш сервис должен имплементировать оба этих метода
|
||
@app.get("/health")
|
||
async def health() -> dict[str, str]:
|
||
return {"status": "ok"}
|
||
|
||
|
||
@app.post("/index", response_model=IndexAPIResponse)
|
||
async def index(payload: IndexAPIRequest) -> IndexAPIResponse:
|
||
return IndexAPIResponse(
|
||
results=build_chunks(
|
||
payload.data.overlap_messages,
|
||
payload.data.new_messages,
|
||
)
|
||
)
|
||
|
||
|
||
@lru_cache(maxsize=1)
|
||
def get_sparse_model():
|
||
from fastembed import SparseTextEmbedding
|
||
|
||
# можете делать любой вектор, который будет совместим с вашим поиском в Qdrant
|
||
# помните об ограничении времени выполнения вашей работы в тестирующей системе
|
||
logger.info(
|
||
"Loading sparse model %s from cache %s",
|
||
SPARSE_MODEL_NAME,
|
||
FASTEMBED_CACHE_PATH,
|
||
)
|
||
return SparseTextEmbedding(model_name=SPARSE_MODEL_NAME)
|
||
|
||
|
||
def embed_sparse_texts(texts: list[str]) -> list[SparseVector]:
|
||
model = get_sparse_model()
|
||
vectors: list[dict[str, list[int] | list[float]]] = []
|
||
|
||
for item in model.embed(texts):
|
||
vectors.append(
|
||
{
|
||
"indices": item.indices.tolist(),
|
||
"values": item.values.tolist(),
|
||
}
|
||
)
|
||
|
||
return vectors
|
||
|
||
|
||
@app.post("/sparse_embedding")
|
||
async def sparse_embedding(payload: SparseEmbeddingRequest) -> dict[str, Any]:
|
||
# Проверяющая система вызывает этот endpoint при создании коллекции
|
||
vectors = await asyncio.to_thread(embed_sparse_texts, payload.texts)
|
||
return {"vectors": vectors}
|
||
|
||
# красивая обработка ошибок
|
||
@app.exception_handler(Exception)
|
||
async def exception_handler(request: Request, exc: Exception) -> JSONResponse:
|
||
logger.exception(exc)
|
||
|
||
if isinstance(exc, RequestValidationError):
|
||
return JSONResponse(status_code=422, content={"detail": exc.errors()})
|
||
|
||
return JSONResponse(status_code=500, content={"detail": str(exc)})
|
||
|
||
|
||
def main() -> None:
|
||
import uvicorn
|
||
|
||
uvicorn.run(
|
||
"main:app",
|
||
host=HOST,
|
||
port=PORT,
|
||
reload=False,
|
||
workers=UVICORN_WORKERS,
|
||
)
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|