import asyncio from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from app.models.rulebook_chunk import RulebookChunk from app.rag.embeddings import embed_query DEFAULT_TOP_K = 5 # Cosine distance cutoff (0 = identical, 2 = opposite) — drop weak matches rather than # stuffing the prompt with irrelevant rulebook text when nothing actually fits the query. # Calibrated against test queries: genuinely relevant chunks scored 0.30-0.45. MAX_DISTANCE = 0.6 # Hard cap on injected rulebook text so one turn's context/cost doesn't balloon even when # several chunks all clear the relevance threshold. MAX_BLOCK_CHARS = 4000 async def retrieve_relevant_chunks( session: AsyncSession, query: str, top_k: int = DEFAULT_TOP_K ) -> list[tuple[RulebookChunk, float]]: # embed_query is a synchronous, CPU-bound sentence-transformers call — run it off the # event loop so it doesn't stall other concurrent requests/WebSocket connections. query_embedding = await asyncio.to_thread(embed_query, query) distance = RulebookChunk.embedding.cosine_distance(query_embedding) rows = ( await session.execute( select(RulebookChunk, distance.label("distance")).order_by(distance).limit(top_k) ) ).all() return [(row[0], row[1]) for row in rows] async def build_rag_block(session: AsyncSession, query: str, top_k: int = DEFAULT_TOP_K) -> str: results = await retrieve_relevant_chunks(session, query, top_k=top_k) relevant = [chunk for chunk, distance in results if distance <= MAX_DISTANCE] if not relevant: return "" parts = [ "Relevante Auszüge aus den Regelwerken — richte dich danach und nenne bei Bedarf die Quelle:" ] used_chars = 0 for chunk in relevant: entry = f"\n[{chunk.source_document}]\n{chunk.content}" if used_chars + len(entry) > MAX_BLOCK_CHARS and used_chars > 0: break parts.append(entry) used_chars += len(entry) return "\n".join(parts)