From 8fd0bd76f9ee2df897d26a3f41b2835879db51b9 Mon Sep 17 00:00:00 2001 From: Roland Tannous Date: Sat, 30 May 2026 14:27:05 +0400 Subject: [PATCH] Revert "Studio: RAG for external model providers via prefetch" This reverts commit b836c3c76be37f5d2ca9c37c8419a849857bce27. --- studio/backend/core/rag/query_decompose.py | 122 ---------- studio/backend/routes/rag.py | 212 +++++------------- .../src/components/assistant-ui/thread.tsx | 9 +- .../src/features/chat/api/chat-adapter.ts | 161 ++++--------- .../src/features/chat/shared-composer.tsx | 11 +- .../frontend/src/features/rag/api/rag-api.ts | 30 --- 6 files changed, 111 insertions(+), 434 deletions(-) delete mode 100644 studio/backend/core/rag/query_decompose.py diff --git a/studio/backend/core/rag/query_decompose.py b/studio/backend/core/rag/query_decompose.py deleted file mode 100644 index 6602827fd6..0000000000 --- a/studio/backend/core/rag/query_decompose.py +++ /dev/null @@ -1,122 +0,0 @@ -# SPDX-License-Identifier: AGPL-3.0-only -# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 - -"""Query decomposition for external-provider RAG prefetch. - -External providers can't run studio's `search_knowledge_base` tool loop, so -for the prefetch path we retrieve up-front. To match the multi-query -behaviour local models get from the tool-loop system prompt, we spin up the -pre-cached helper GGUF (``unsloth/gemma-4-E2B-it-GGUF``, the same model the -captioner uses) momentarily, ask it to split the user's question into up to -three focused search queries, then unload it. - -The helper llama-server is its own subprocess (llama.cpp), spawned with -``kill_orphans=False`` so it can't reap a resident chat model, and always -unloaded in a ``finally``. Any failure (helper can't load, request errors, -empty output) falls back to ``[query]`` — a single raw retrieval — so RAG -prefetch never hard-fails on decomposition. -""" - -from __future__ import annotations - -from typing import Any, Optional - -from loggers import get_logger - -# Reuse the exact model the captioner / precache path already downloads. -from core.rag.captioner import _HELPER_REPO, _HELPER_VARIANT, _HELPER_MODEL_NAME - -logger = get_logger(__name__) - -_MAX_QUERIES = 3 -_REQUEST_TIMEOUT_SECONDS = 60.0 -_PROMPT = ( - "Split the user's question into up to 3 focused search queries for " - "retrieving relevant passages from their documents. Prefer fewer when " - "the question is narrow — one is fine. Output ONLY the queries, one per " - "line, no numbering, no preamble." -) - - -def _load_helper() -> Optional[tuple[Any, str, str]]: - """Spawn a private text-only helper llama-server. Caller unloads it.""" - try: - from core.inference.llama_cpp import LlamaCppBackend - - # kill_orphans=False: a resident chat-model llama-server (if any) - # must not be reaped by this transient instance. - backend = LlamaCppBackend(kill_orphans = False) - ok = backend.load_model( - hf_repo = _HELPER_REPO, - hf_variant = _HELPER_VARIANT, - model_identifier = f"rag-querygen:{_HELPER_REPO}:{_HELPER_VARIANT}", - is_vision = False, - n_ctx = 4096, - n_gpu_layers = -1, - ) - if not ok: - logger.warning("RAG query-decompose: helper failed to start") - return None - return backend, backend.base_url, _HELPER_MODEL_NAME - except Exception as exc: # noqa: BLE001 - logger.warning("RAG query-decompose: helper load raised", error = str(exc)) - return None - - -def _parse_queries(raw: str, fallback: str) -> list[str]: - out: list[str] = [] - for line in (raw or "").splitlines(): - # Strip common list markers the model might emit despite the prompt. - cleaned = line.strip().lstrip("-*0123456789.) ").strip() - if cleaned: - out.append(cleaned) - if len(out) >= _MAX_QUERIES: - break - return out or [fallback] - - -def decompose_query(query: str) -> list[str]: - """Return up to 3 focused search queries; ``[query]`` on any failure. - - Loads the helper, asks for the decomposition, unloads. Never raises. - """ - q = (query or "").strip() - if not q: - return [] - - import httpx - - loaded = _load_helper() - if loaded is None: - return [q] - backend, base_url, model_name = loaded - try: - endpoint = f"{base_url.rstrip('/')}/v1/chat/completions" - payload = { - "model": model_name, - "messages": [ - {"role": "system", "content": _PROMPT}, - {"role": "user", "content": q}, - ], - "max_tokens": 160, - "temperature": 0.0, - # gemma-4 is a reasoning model; thinking would eat the budget and - # emit no visible queries (same issue the captioner hit). - "chat_template_kwargs": {"enable_thinking": False}, - } - with httpx.Client(timeout = _REQUEST_TIMEOUT_SECONDS) as client: - response = client.post(endpoint, json = payload) - response.raise_for_status() - data = response.json() - content = data.get("choices", [{}])[0].get("message", {}).get("content", "") - queries = _parse_queries(content if isinstance(content, str) else "", q) - logger.info("RAG query-decompose: produced queries", n = len(queries)) - return queries - except Exception as exc: # noqa: BLE001 - logger.warning("RAG query-decompose: request failed", error = str(exc)) - return [q] - finally: - try: - backend.unload_model() - except Exception as exc: # noqa: BLE001 - logger.warning("RAG query-decompose: helper unload failed", error = str(exc)) diff --git a/studio/backend/routes/rag.py b/studio/backend/routes/rag.py index 0d8aca6ce3..512f5de3bb 100644 --- a/studio/backend/routes/rag.py +++ b/studio/backend/routes/rag.py @@ -160,26 +160,6 @@ class SearchResponse(BaseModel): hits: list[SearchHit] -class PrefetchRequest(BaseModel): - """Prefetch RAG for an external-provider turn: studio decomposes the - question (via the pre-cached helper) and retrieves, so the frontend can - inject the chunks before calling the provider.""" - - query: str = Field(min_length = 1, max_length = 4000) - kb_id: str | None = None - thread_id: str | None = None - top_k: int = Field(default = 10, ge = 1, le = 100) - mode: Literal["bm25", "dense", "hybrid"] = "hybrid" - enable_rerank: bool = False - reranker_model: str | None = None - min_score: float = Field(default = 0.0, ge = 0.0, le = 1.0) - - -class PrefetchResponse(BaseModel): - queries: list[str] - hits: list[SearchHit] - - # --- Helpers --- @@ -1650,43 +1630,72 @@ def get_document_file( ) -def _execute_search( - scope: str, - *, - scope_embedder: str | None, - query: str, - mode: str, - top_k: int, - document_ids: list[str] | None, - enable_rerank: bool, - reranker_model: str | None, - min_score: float, -) -> list[SearchHit]: - """Run one retrieval against an already-resolved scope. Shared by the - /search and /prefetch endpoints.""" - candidate_k = max(top_k, RAG_RERANK_CANDIDATE_K) if enable_rerank else top_k +@router.post("/search", response_model = SearchResponse) +def search( + payload: SearchRequest, + current_subject: str = Depends(get_current_subject), +) -> SearchResponse: + if bool(payload.kb_id) == bool(payload.thread_id): + raise HTTPException( + status_code = 400, + detail = "exactly one of kb_id or thread_id must be supplied", + ) + if payload.kb_id: + _kb_or_404(payload.kb_id) + scope = kb_scope(payload.kb_id) + else: + scope = thread_scope(payload.thread_id) - if mode == "bm25": - hits = retrieval.retrieve_bm25(scope, query, candidate_k) - elif mode == "dense": + # Query must use the same embedder as the scope (dim must match). + scope_embedder = _resolve_scope_embedder(scope) + logger.info( + "RAG search: scope=%s embedder=%s mode=%s top_k=%d min_score=%.3f rerank=%s query=%r", + scope, + scope_embedder or "", + payload.mode, + payload.top_k, + payload.min_score, + payload.enable_rerank, + payload.query[:120], + ) + + # Reranker needs a wider candidate pool than top_k. + candidate_k = ( + max(payload.top_k, RAG_RERANK_CANDIDATE_K) + if payload.enable_rerank + else payload.top_k + ) + + if payload.mode == "bm25": + hits = retrieval.retrieve_bm25(scope, payload.query, candidate_k) + elif payload.mode == "dense": hits = retrieval.retrieve_dense( scope, - query, + payload.query, candidate_k, - document_ids = document_ids, + document_ids = payload.document_ids, embedder_model = scope_embedder, ) else: hits = retrieval.retrieve_hybrid( scope, - query, + payload.query, k = candidate_k, - document_ids = document_ids, + document_ids = payload.document_ids, embedder_model = scope_embedder, ) - if min_score > 0.0: - hits = retrieval.filter_by_min_score(hits, min_score) + retrieved_count = len(hits) + if payload.min_score > 0.0: + hits = retrieval.filter_by_min_score(hits, payload.min_score) + logger.info( + "RAG search: retrieved=%d met_threshold=%d (min_score=%.3f)", + retrieved_count, + len(hits), + payload.min_score, + ) + else: + logger.info("RAG search: retrieved=%d (no threshold)", retrieved_count) chunk_ids = [h.chunk_id for h in hits] chunk_lookup: dict[str, dict] = {} @@ -1709,20 +1718,20 @@ def _execute_search( for r in rows: chunk_lookup[r["chunk_id"]] = dict(r) - if enable_rerank: + if payload.enable_rerank: pairs = [ (hit, chunk_lookup[hit.chunk_id]["text"]) for hit in hits if hit.chunk_id in chunk_lookup ] hits = reranker.rerank( - query, + payload.query, pairs, - model_name = reranker_model, - top_k = top_k, + model_name = payload.reranker_model, + top_k = payload.top_k, ) else: - hits = hits[:top_k] + hits = hits[: payload.top_k] out: list[SearchHit] = [] for hit in hits: @@ -1753,110 +1762,5 @@ def _execute_search( line_end = meta.get("line_end"), ) ) - return out - - -@router.post("/search", response_model = SearchResponse) -def search( - payload: SearchRequest, - current_subject: str = Depends(get_current_subject), -) -> SearchResponse: - if bool(payload.kb_id) == bool(payload.thread_id): - raise HTTPException( - status_code = 400, - detail = "exactly one of kb_id or thread_id must be supplied", - ) - if payload.kb_id: - _kb_or_404(payload.kb_id) - scope = kb_scope(payload.kb_id) - else: - scope = thread_scope(payload.thread_id) - - # Query must use the same embedder as the scope (dim must match). - scope_embedder = _resolve_scope_embedder(scope) - logger.info( - "RAG search: scope=%s embedder=%s mode=%s top_k=%d min_score=%.3f rerank=%s query=%r", - scope, - scope_embedder or "", - payload.mode, - payload.top_k, - payload.min_score, - payload.enable_rerank, - payload.query[:120], - ) - - out = _execute_search( - scope, - scope_embedder = scope_embedder, - query = payload.query, - mode = payload.mode, - top_k = payload.top_k, - document_ids = payload.document_ids, - enable_rerank = payload.enable_rerank, - reranker_model = payload.reranker_model, - min_score = payload.min_score, - ) logger.info("RAG search: returning %d hits", len(out)) return SearchResponse(hits = out) - - -@router.post("/prefetch", response_model = PrefetchResponse) -def prefetch( - payload: PrefetchRequest, - current_subject: str = Depends(get_current_subject), -) -> PrefetchResponse: - """External-provider RAG prefetch: decompose the question via the helper, - retrieve per sub-query, merge+dedup, return chunks for the frontend to - inject before calling the provider. The local tool path is unaffected.""" - if bool(payload.kb_id) == bool(payload.thread_id): - raise HTTPException( - status_code = 400, - detail = "exactly one of kb_id or thread_id must be supplied", - ) - if payload.kb_id: - _kb_or_404(payload.kb_id) - scope = kb_scope(payload.kb_id) - else: - scope = thread_scope(payload.thread_id) - - scope_embedder = _resolve_scope_embedder(scope) - - # Momentarily load the pre-cached helper to split the question into up to - # 3 focused queries; falls back to [query] on any failure. - from core.rag.query_decompose import decompose_query - - queries = decompose_query(payload.query) - logger.info( - "RAG prefetch: scope=%s embedder=%s mode=%s n_queries=%d rerank=%s", - scope, - scope_embedder or "", - payload.mode, - len(queries), - payload.enable_rerank, - ) - - # Retrieve per query, merge, dedup by chunk_id (keep first/highest-ranked - # occurrence), then cap at top_k. - merged: list[SearchHit] = [] - seen: set[str] = set() - for q in queries: - hits = _execute_search( - scope, - scope_embedder = scope_embedder, - query = q, - mode = payload.mode, - top_k = payload.top_k, - document_ids = None, - enable_rerank = payload.enable_rerank, - reranker_model = payload.reranker_model, - min_score = payload.min_score, - ) - for h in hits: - if h.chunk_id in seen: - continue - seen.add(h.chunk_id) - merged.append(h) - - merged = merged[: payload.top_k] - logger.info("RAG prefetch: returning %d merged hits", len(merged)) - return PrefetchResponse(queries = queries, hits = merged) diff --git a/studio/frontend/src/components/assistant-ui/thread.tsx b/studio/frontend/src/components/assistant-ui/thread.tsx index acdab53539..b09d836692 100644 --- a/studio/frontend/src/components/assistant-ui/thread.tsx +++ b/studio/frontend/src/components/assistant-ui/thread.tsx @@ -1151,12 +1151,9 @@ const RagToggle: FC = () => { const ragSource = useChatRuntimeStore((s) => s.ragSource); const setRagSource = useChatRuntimeStore((s) => s.setRagSource); const supportsTools = useChatRuntimeStore((s) => s.supportsTools); - const checkpoint = useChatRuntimeStore((s) => s.params.checkpoint); - const isExternalModel = parseExternalModelId(checkpoint) !== null; - // Local models need tool-calling (search_knowledge_base loop); external - // providers use the prefetch path, so RAG is allowed regardless of the - // local supportsTools flag (mirrors shared-composer's ragDisabled). - const disabled = !modelLoaded || (!supportsTools && !isExternalModel); + // RAG runs through the local search_knowledge_base tool, so it needs + // tool-calling support (mirrors shared-composer's ragDisabled). + const disabled = !modelLoaded || !supportsTools; return (