vk_hackathon/index/main.py
2026-04-18 12:38:29 +03:00

687 lines
21 KiB
Python
Raw Permalink 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.

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()