From 3b477a816c6cc5f8f8e3adc0d7e884f17e06c8b1 Mon Sep 17 00:00:00 2001 From: Roland Tannous Date: Tue, 26 May 2026 12:39:50 +0400 Subject: [PATCH] Studio: add BM25/semantic/hybrid search-mode toggle to RAG settings --- studio/backend/core/inference/tools.py | 5 +++- studio/backend/core/rag/tool.py | 28 ++++++++++++----- studio/backend/models/inference.py | 3 +- .../src/features/chat/api/chat-adapter.ts | 7 +++-- .../features/chat/api/chat-settings-api.ts | 3 ++ .../src/features/chat/chat-settings-sheet.tsx | 30 ++++++++++++++++++- .../chat/stores/chat-runtime-store.ts | 12 +++++++- 7 files changed, 74 insertions(+), 14 deletions(-) diff --git a/studio/backend/core/inference/tools.py b/studio/backend/core/inference/tools.py index 0d277bb65b..085235e8fb 100644 --- a/studio/backend/core/inference/tools.py +++ b/studio/backend/core/inference/tools.py @@ -538,7 +538,7 @@ def execute_tool( ``tool_context``: optional per-request extras the LLM does not see (RAG scope, future per-tool overrides). Keys consumed: - ``rag_scope``: ``{kb_id?, thread_id?, enable_rerank?, default_top_k?, - reranker_model?}`` — consumed by ``search_knowledge_base``. + reranker_model?, min_score?, mode?}`` — consumed by ``search_knowledge_base``. """ logger.info( f"execute_tool: name={name}, session_id={session_id}, timeout={timeout}" @@ -562,6 +562,8 @@ def execute_tool( from core.rag.tool import search_knowledge_base scope = (tool_context or {}).get("rag_scope") or {} + raw_mode = scope.get("mode") + mode = raw_mode if raw_mode in ("bm25", "dense", "hybrid") else "hybrid" return search_knowledge_base( query = arguments.get("query", ""), top_k = arguments.get("top_k"), @@ -571,6 +573,7 @@ def execute_tool( reranker_model = scope.get("reranker_model"), default_top_k = int(scope.get("default_top_k") or 5), min_score = float(scope.get("min_score") or 0.0), + mode = mode, ) return f"Unknown tool: {name}" diff --git a/studio/backend/core/rag/tool.py b/studio/backend/core/rag/tool.py index 9156ca8324..178da09991 100644 --- a/studio/backend/core/rag/tool.py +++ b/studio/backend/core/rag/tool.py @@ -16,7 +16,7 @@ LLM doesn't need to know about KB UUIDs. from __future__ import annotations -from typing import Any +from typing import Any, Literal from loggers import get_logger @@ -93,6 +93,7 @@ def search_knowledge_base( reranker_model: str | None = None, default_top_k: int = 5, min_score: float = 0.0, + mode: Literal["bm25", "dense", "hybrid"] = "hybrid", ) -> str: """Execute the RAG search and return a tool-result string. @@ -130,9 +131,10 @@ def search_knowledge_base( scope_embedder = resolve_scope_embedder(scope) logger.info( - "search_knowledge_base: scope=%s embedder=%s top_k=%d min_score=%.3f rerank=%s query=%r", + "search_knowledge_base: scope=%s embedder=%s mode=%s top_k=%d min_score=%.3f rerank=%s query=%r", scope, scope_embedder or "", + mode, k, min_score, enable_rerank, @@ -140,12 +142,22 @@ def search_knowledge_base( ) try: - hits = retrieval.retrieve_hybrid( - scope, - query.strip(), - k = candidate_k, - embedder_model = scope_embedder, - ) + if mode == "bm25": + hits = retrieval.retrieve_bm25(scope, query.strip(), candidate_k) + elif mode == "dense": + hits = retrieval.retrieve_dense( + scope, + query.strip(), + candidate_k, + embedder_model = scope_embedder, + ) + else: + hits = retrieval.retrieve_hybrid( + scope, + query.strip(), + k = candidate_k, + embedder_model = scope_embedder, + ) except Exception as exc: # noqa: BLE001 logger.exception("search_knowledge_base retrieval failed") return f"Error: retrieval failed ({type(exc).__name__})." diff --git a/studio/backend/models/inference.py b/studio/backend/models/inference.py index 091dba4f8d..b4d6bb926b 100644 --- a/studio/backend/models/inference.py +++ b/studio/backend/models/inference.py @@ -692,7 +692,8 @@ class ChatCompletionRequest(BaseModel): "[x-unsloth] Per-request context the `search_knowledge_base` tool " "consumes when the LLM invokes it. Shape: " "{kb_id?: str, thread_id?: str, enable_rerank?: bool, " - "default_top_k?: int, reranker_model?: str}. Ignored unless " + "default_top_k?: int, reranker_model?: str, min_score?: float, " + "mode?: 'bm25'|'dense'|'hybrid'}. Ignored unless " "'search_knowledge_base' is in enabled_tools." ), ) diff --git a/studio/frontend/src/features/chat/api/chat-adapter.ts b/studio/frontend/src/features/chat/api/chat-adapter.ts index aebf3d4f61..b743948012 100644 --- a/studio/frontend/src/features/chat/api/chat-adapter.ts +++ b/studio/frontend/src/features/chat/api/chat-adapter.ts @@ -61,7 +61,7 @@ import { type SearchRequest, search as ragSearch, } from "@/features/rag/api/rag-api"; -import type { RagSource } from "./chat-settings-api"; +import type { RagMode, RagSource } from "./chat-settings-api"; import { createOpenAIContainer, listOpenAIContainers, @@ -95,11 +95,12 @@ function buildRagRequest( enableRerank: boolean, topK: number, minScore: number, + mode: RagMode, ): SearchRequest | null { const base: SearchRequest = { query, top_k: topK, - mode: "hybrid", + mode, enable_rerank: enableRerank, min_score: minScore, }; @@ -1023,6 +1024,7 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter { runtime.enableRerank, runtime.ragTopK, runtime.ragMinScore, + runtime.ragMode, ); if (ragReq) { try { @@ -1664,6 +1666,7 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter { enable_rerank: runtime.enableRerank, default_top_k: runtime.ragTopK, min_score: runtime.ragMinScore, + mode: runtime.ragMode, }, } : {}), diff --git a/studio/frontend/src/features/chat/api/chat-settings-api.ts b/studio/frontend/src/features/chat/api/chat-settings-api.ts index c85d8af3da..2b7b0f8aa3 100644 --- a/studio/frontend/src/features/chat/api/chat-settings-api.ts +++ b/studio/frontend/src/features/chat/api/chat-settings-api.ts @@ -20,6 +20,8 @@ export type RagSource = | { kind: "thread" } | { kind: "kb"; kbId: string }; +export type RagMode = "bm25" | "dense" | "hybrid"; + export interface PersistedChatSettings { inferenceParams?: PersistedInferenceParams; customPresets?: PersistedChatPreset[]; @@ -32,6 +34,7 @@ export interface PersistedChatSettings { maxToolCallsPerMessage?: number; toolCallTimeout?: number; ragSource?: RagSource; + ragMode?: RagMode; enableRerank?: boolean; ragTopK?: number; ragMinScore?: number; diff --git a/studio/frontend/src/features/chat/chat-settings-sheet.tsx b/studio/frontend/src/features/chat/chat-settings-sheet.tsx index 756e43fa7f..3f3bdef19a 100644 --- a/studio/frontend/src/features/chat/chat-settings-sheet.tsx +++ b/studio/frontend/src/features/chat/chat-settings-sheet.tsx @@ -90,7 +90,7 @@ import { } from "./provider-capabilities"; import { useChatRuntimeStore } from "./stores/chat-runtime-store"; import type { InferenceParams } from "./types/runtime"; -import type { RagSource } from "./api/chat-settings-api"; +import type { RagMode, RagSource } from "./api/chat-settings-api"; import { DocumentRow } from "@/features/rag/components/document-row"; import { KBCreateDialog } from "@/features/rag/components/kb-create-dialog"; import { useKnowledgeBases } from "@/features/rag/hooks/use-knowledge-bases"; @@ -441,6 +441,8 @@ export function ChatSettingsPanel({ const isGguf = useChatRuntimeStore((s) => s.activeGgufVariant) != null; const ragSource = useChatRuntimeStore((s) => s.ragSource); const setRagSource = useChatRuntimeStore((s) => s.setRagSource); + const ragMode = useChatRuntimeStore((s) => s.ragMode); + const setRagMode = useChatRuntimeStore((s) => s.setRagMode); const ragToolEnabled = useChatRuntimeStore((s) => s.ragToolEnabled); const enableRerank = useChatRuntimeStore((s) => s.enableRerank); const setEnableRerank = useChatRuntimeStore((s) => s.setEnableRerank); @@ -1360,6 +1362,32 @@ export function ChatSettingsPanel({ source before sending.

+
+ + +

+ Hybrid blends keyword (BM25) with vector similarity — best + default. Semantic-only ignores exact terms; BM25-only ignores + meaning. +

+
void; setContextUsage: (usage: ChatRuntimeStore["contextUsage"]) => void; setRagSource: (source: RagSource) => void; + setRagMode: (mode: RagMode) => void; setEnableRerank: (value: boolean) => void; setRagTopK: (value: number) => void; setRagMinScore: (value: number) => void; @@ -370,6 +372,7 @@ type ScalarSettingKey = | "maxToolCallsPerMessage" | "toolCallTimeout" | "ragSource" + | "ragMode" | "enableRerank" | "ragTopK" | "ragMinScore"; @@ -407,6 +410,7 @@ const SCALAR_SETTING_KEYS = [ "maxToolCallsPerMessage", "toolCallTimeout", "ragSource", + "ragMode", "enableRerank", "ragTopK", "ragMinScore", @@ -619,6 +623,7 @@ export const useChatRuntimeStore = create((set, get) => ({ modelLoading: false, activeNativePathToken: null, ragSource: { kind: "thread" }, + ragMode: "hybrid", enableRerank: false, ragTopK: 5, ragMinScore: 0, @@ -841,6 +846,11 @@ export const useChatRuntimeStore = create((set, get) => ({ setScalarSettingVersion("ragSource", ragSource, state.ragSource); return { ragSource }; }), + setRagMode: (ragMode) => + set((state) => { + setScalarSettingVersion("ragMode", ragMode, state.ragMode); + return { ragMode }; + }), setEnableRerank: (enableRerank) => set((state) => { setScalarSettingVersion(