Studio: add BM25/semantic/hybrid search-mode toggle to RAG settings

This commit is contained in:
Roland Tannous 2026-05-26 12:39:50 +04:00
commit 3b477a816c
7 changed files with 74 additions and 14 deletions

View file

@ -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}"

View file

@ -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 "<default>",
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__})."

View file

@ -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."
),
)

View file

@ -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,
},
}
: {}),

View file

@ -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;

View file

@ -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.
</p>
</div>
<div className="flex flex-col gap-1.5">
<label className="text-[12px] font-medium text-muted-foreground">
Search mode
</label>
<Select
value={ragMode}
onValueChange={(v) => setRagMode(v as RagMode)}
disabled={!ragEnabled}
>
<SelectTrigger className="w-full">
<SelectValue />
</SelectTrigger>
<SelectContent>
<SelectItem value="hybrid">
Hybrid (BM25 + semantic)
</SelectItem>
<SelectItem value="dense">Semantic only</SelectItem>
<SelectItem value="bm25">BM25 (lexical) only</SelectItem>
</SelectContent>
</Select>
<p className="text-[11px] text-muted-foreground">
Hybrid blends keyword (BM25) with vector similarity best
default. Semantic-only ignores exact terms; BM25-only ignores
meaning.
</p>
</div>
<KBCreateDialog
open={kbCreateOpen}
onOpenChange={setKbCreateOpen}

View file

@ -19,7 +19,7 @@ import {
loadChatSettingsWithLegacyImport,
savePersistedChatSettingsPatch,
} from "../utils/chat-settings-storage";
import type { RagSource } from "../api/chat-settings-api";
import type { RagMode, RagSource } from "../api/chat-settings-api";
const HF_TOKEN_KEY = "unsloth_hf_token";
export const CHAT_REASONING_ENABLED_KEY = "unsloth_chat_reasoning_enabled";
@ -298,6 +298,7 @@ type ChatRuntimeStore = {
modelLoading: boolean;
activeNativePathToken: string | null;
ragSource: RagSource;
ragMode: RagMode;
enableRerank: boolean;
ragTopK: number;
// Cosine-similarity floor for RAG hits. Off (0) by default — set
@ -349,6 +350,7 @@ type ChatRuntimeStore = {
clearPendingAudio: () => 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<ChatRuntimeStore>((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<ChatRuntimeStore>((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(