unsloth/studio/backend/core/rag/tool.py
Daniel Han 8848a310df
Studio: clean-room compact RAG (knowledge bases, hybrid search, fast indexing) (#5910)
Adds a self-contained RAG stack to Studio: knowledge bases with chunked indexing, hybrid (dense + lexical) retrieval, and an automatic first-pass context inject into chat. Embeddings run through a local llama-server GGUF backend (default unsloth/bge-small-en-v1.5-GGUF) with a sentence-transformers fallback. The chat tool loop gains a search_knowledge_base tool, a per-turn re-search cap, and source citation, layered on top of the shared ToolLoopController.
2026-06-09 21:17:04 -07:00

193 lines
6 KiB
Python

# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
"""``search_knowledge_base`` LLM tool: scope resolution + hit formatting.
KB scope wins over thread scope. Hits render as ``<chunk>`` blocks for the model,
plus a parallel citation source-map for clickable sources. Each call opens and
closes its own ``rag_db`` connection.
"""
from __future__ import annotations
from xml.sax.saxutils import quoteattr
from storage import rag_db
from . import config, retrieval
from .store import kb_scope, thread_scope
SEARCH_KNOWLEDGE_BASE_TOOL = {
"type": "function",
"function": {
"name": "search_knowledge_base",
"description": (
"Search the user's uploaded documents and knowledge bases for relevant passages."
),
"parameters": {
"type": "object",
"properties": {
"query": {
"type": "string",
"description": "Natural-language search query.",
},
"top_k": {
"type": "integer",
"description": "Max chunks to return.",
},
},
"required": ["query"],
},
},
}
def _resolve_scope(scope_kb_id: str | None, scope_thread_id: str | None) -> str | None:
if scope_kb_id:
return kb_scope(scope_kb_id)
if scope_thread_id:
return thread_scope(scope_thread_id)
return None
def _format(rows, hits) -> tuple[str, list[dict]]:
"""Render hits as ``<chunk>`` blocks and build a citation source-map."""
if not hits:
return "No matching chunks were found in the knowledge base.", []
blocks: list[str] = []
sources: list[dict] = []
for i, h in enumerate(hits, 1):
r = rows.get(h.chunk_id)
filename = (r["filename"] if r else None) or "unknown"
page = r["page_number"] if r else None
text = r["text"] if r else ""
src = quoteattr(filename)
page_attr = f" page={quoteattr(str(page))}" if page else ""
blocks.append(f'<chunk id="{i}" source={src}{page_attr}>\n{text}\n</chunk>')
sources.append(
{
"citationId": i,
"chunkId": h.chunk_id,
"documentId": r["document_id"] if r else None,
"filename": filename,
"page": page,
"text": text,
"score": round(float(h.score), 4) if h.score is not None else None,
}
)
return "\n\n".join(blocks), sources
def search_knowledge_base_with_sources(
*,
query: str,
scope_kb_id: str | None = None,
scope_thread_id: str | None = None,
top_k: int | None = None,
min_score: float = 0.0,
model_name: str | None = None,
mode: str = "hybrid",
) -> tuple[str, list[dict]]:
"""Search -> ``(rendered_text, citation_sources)``; each source aligns with a
rendered ``<chunk>`` block's ``id``."""
if not query or not query.strip():
return "Error: query is empty.", []
scope = _resolve_scope(scope_kb_id, scope_thread_id)
if scope is None:
return "No documents are attached to this chat.", []
conn = rag_db.get_connection()
try:
hits = retrieval.retrieve_hybrid(
conn,
scope,
query,
k = top_k or config.TOP_K_HYBRID,
model_name = model_name,
mode = mode,
)
hits = retrieval.filter_min_score(hits, min_score)
rows = store_rows(conn, hits)
finally:
conn.close()
return _format(rows, hits)
def store_rows(conn, hits):
"""Hydrate chunk rows for a list of hits."""
from . import store
return store.chunks_by_id(conn, [h.chunk_id for h in hits])
def search_for_autoinject(
*,
query: str,
scope_kb_id: str | None = None,
scope_thread_id: str | None = None,
top_k: int | None = None,
min_dense_score: float = 0.70,
model_name: str | None = None,
mode: str = "hybrid",
) -> tuple[str, list[dict]] | None:
"""Forced-retrieval variant for auto-injection.
Returns ``(rendered_text, sources)`` only if some hit's cosine clears
``min_dense_score``, else ``None`` (inject nothing). The dense gate keeps
weak/off-topic matches out of answers. In ``lexical`` mode hits carry no
cosine, so the gate falls back to a dense 1-NN probe.
"""
if not query or not query.strip():
return None
scope = _resolve_scope(scope_kb_id, scope_thread_id)
if scope is None:
return None
k = top_k or config.TOP_K_HYBRID
conn = rag_db.get_connection()
try:
hits = retrieval.retrieve_hybrid(
conn,
scope,
query,
k = k,
model_name = model_name,
mode = mode,
)
strong = [
h for h in hits if h.dense_score is not None and h.dense_score >= min_dense_score
][:k]
if not strong and hits and mode == "lexical":
probe = retrieval.retrieve_dense(conn, scope, query, 1, model_name = model_name)
if (
probe
and probe[0].dense_score is not None
and (probe[0].dense_score >= min_dense_score)
):
strong = hits[:k]
if not strong:
return None
rows = store_rows(conn, strong)
finally:
conn.close()
text, sources = _format(rows, strong)
return (text, sources) if sources else None
def search_knowledge_base(
*,
query: str,
scope_kb_id: str | None = None,
scope_thread_id: str | None = None,
top_k: int | None = None,
min_score: float = 0.0,
model_name: str | None = None,
) -> str:
"""Text-only variant of :func:`search_knowledge_base_with_sources`."""
text, _sources = search_knowledge_base_with_sources(
query = query,
scope_kb_id = scope_kb_id,
scope_thread_id = scope_thread_id,
top_k = top_k,
min_score = min_score,
model_name = model_name,
)
return text