mirror of
https://github.com/RooVetGit/Roo-Code.git
synced 2026-08-28 05:27:24 +00:00
fix: separate document and query prefixes for embedding models
Adds EmbeddingPurpose type and documentPrefix field to support different prefixes for indexing vs querying. Models like nomic-embed-code require a query prefix for search but no prefix for document indexing. - Add documentPrefix to EmbeddingModelProfile type - Add getModelDocumentPrefix() and getModelPrefixForPurpose() helpers - Add purpose parameter to IEmbedder.createEmbeddings() - Update all embedder implementations to use purpose-aware prefixing - Scanner and file-watcher pass "index" purpose - Search service passes "query" purpose - Update nomic-embed-code profiles (no documentPrefix for indexing) - Update tests for new behavior Closes #11707
This commit is contained in:
parent
137d3f4fd8
commit
21e967c35f
18 changed files with 171 additions and 36 deletions
|
|
@ -8,10 +8,13 @@ export type EmbedderProvider =
|
|||
| "bedrock"
|
||||
| "openrouter" // Add other providers as needed.
|
||||
|
||||
export type EmbeddingPurpose = "index" | "query"
|
||||
|
||||
export interface EmbeddingModelProfile {
|
||||
dimension: number
|
||||
scoreThreshold?: number // Model-specific minimum score threshold for semantic search.
|
||||
queryPrefix?: string // Optional prefix required by the model for queries.
|
||||
queryPrefix?: string // Optional prefix required by the model for search queries.
|
||||
documentPrefix?: string // Optional prefix required by the model for indexing/document embedding.
|
||||
// Add other model-specific properties if needed, e.g., context window size.
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -121,7 +121,7 @@ describe("GeminiEmbedder", () => {
|
|||
const result = await embedder.createEmbeddings(texts)
|
||||
|
||||
// Assert
|
||||
expect(mockCreateEmbeddings).toHaveBeenCalledWith(texts, "gemini-embedding-001")
|
||||
expect(mockCreateEmbeddings).toHaveBeenCalledWith(texts, "gemini-embedding-001", undefined)
|
||||
expect(result).toEqual(mockResponse)
|
||||
})
|
||||
|
||||
|
|
@ -141,7 +141,7 @@ describe("GeminiEmbedder", () => {
|
|||
const result = await embedder.createEmbeddings(texts, "gemini-embedding-001")
|
||||
|
||||
// Assert
|
||||
expect(mockCreateEmbeddings).toHaveBeenCalledWith(texts, "gemini-embedding-001")
|
||||
expect(mockCreateEmbeddings).toHaveBeenCalledWith(texts, "gemini-embedding-001", undefined)
|
||||
expect(result).toEqual(mockResponse)
|
||||
})
|
||||
|
||||
|
|
|
|||
|
|
@ -104,7 +104,7 @@ describe("MistralEmbedder", () => {
|
|||
const result = await embedder.createEmbeddings(texts)
|
||||
|
||||
// Assert
|
||||
expect(mockCreateEmbeddings).toHaveBeenCalledWith(texts, "codestral-embed-2505")
|
||||
expect(mockCreateEmbeddings).toHaveBeenCalledWith(texts, "codestral-embed-2505", undefined)
|
||||
expect(result).toEqual(mockResponse)
|
||||
})
|
||||
|
||||
|
|
@ -124,7 +124,7 @@ describe("MistralEmbedder", () => {
|
|||
const result = await embedder.createEmbeddings(texts, "codestral-embed-2505")
|
||||
|
||||
// Assert
|
||||
expect(mockCreateEmbeddings).toHaveBeenCalledWith(texts, "codestral-embed-2505")
|
||||
expect(mockCreateEmbeddings).toHaveBeenCalledWith(texts, "codestral-embed-2505", undefined)
|
||||
expect(result).toEqual(mockResponse)
|
||||
})
|
||||
|
||||
|
|
|
|||
|
|
@ -95,6 +95,7 @@ describe("VercelAiGatewayEmbedder", () => {
|
|||
expect(mockOpenAICompatibleEmbedder.createEmbeddings).toHaveBeenCalledWith(
|
||||
texts,
|
||||
"openai/text-embedding-3-large",
|
||||
undefined,
|
||||
)
|
||||
expect(result).toBe(expectedResponse)
|
||||
})
|
||||
|
|
@ -110,7 +111,7 @@ describe("VercelAiGatewayEmbedder", () => {
|
|||
const result = await embedder.createEmbeddings(texts, customModel)
|
||||
|
||||
// Assert
|
||||
expect(mockOpenAICompatibleEmbedder.createEmbeddings).toHaveBeenCalledWith(texts, customModel)
|
||||
expect(mockOpenAICompatibleEmbedder.createEmbeddings).toHaveBeenCalledWith(texts, customModel, undefined)
|
||||
expect(result).toBe(expectedResponse)
|
||||
})
|
||||
|
||||
|
|
@ -125,6 +126,7 @@ describe("VercelAiGatewayEmbedder", () => {
|
|||
expect(mockOpenAICompatibleEmbedder.createEmbeddings).toHaveBeenCalledWith(
|
||||
texts,
|
||||
"openai/text-embedding-3-large",
|
||||
undefined,
|
||||
)
|
||||
})
|
||||
})
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import type { EmbeddingPurpose } from "@roo-code/types"
|
||||
import { BedrockRuntimeClient, InvokeModelCommand, InvokeModelCommandInput } from "@aws-sdk/client-bedrock-runtime"
|
||||
import { fromIni, fromNodeProviderChain } from "@aws-sdk/credential-providers"
|
||||
import { IEmbedder, EmbeddingResponse, EmbedderInfo } from "../interfaces"
|
||||
|
|
@ -55,7 +56,7 @@ export class BedrockEmbedder implements IEmbedder {
|
|||
* @param model Optional model identifier
|
||||
* @returns Promise resolving to embedding response
|
||||
*/
|
||||
async createEmbeddings(texts: string[], model?: string): Promise<EmbeddingResponse> {
|
||||
async createEmbeddings(texts: string[], model?: string, _purpose?: EmbeddingPurpose): Promise<EmbeddingResponse> {
|
||||
const modelToUse = model || this.defaultModelId
|
||||
|
||||
const allEmbeddings: number[][] = []
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import type { EmbeddingPurpose } from "@roo-code/types"
|
||||
import { OpenAICompatibleEmbedder } from "./openai-compatible"
|
||||
import { IEmbedder, EmbeddingResponse, EmbedderInfo } from "../interfaces/embedder"
|
||||
import { GEMINI_MAX_ITEM_TOKENS } from "../constants"
|
||||
|
|
@ -68,11 +69,11 @@ export class GeminiEmbedder implements IEmbedder {
|
|||
* @param model Optional model identifier (uses constructor model if not provided)
|
||||
* @returns Promise resolving to embedding response
|
||||
*/
|
||||
async createEmbeddings(texts: string[], model?: string): Promise<EmbeddingResponse> {
|
||||
async createEmbeddings(texts: string[], model?: string, purpose?: EmbeddingPurpose): Promise<EmbeddingResponse> {
|
||||
try {
|
||||
// Use the provided model or fall back to the instance's model
|
||||
const modelToUse = model || this.modelId
|
||||
return await this.openAICompatibleEmbedder.createEmbeddings(texts, modelToUse)
|
||||
return await this.openAICompatibleEmbedder.createEmbeddings(texts, modelToUse, purpose)
|
||||
} catch (error) {
|
||||
TelemetryService.instance.captureEvent(TelemetryEventName.CODE_INDEX_ERROR, {
|
||||
error: error instanceof Error ? error.message : String(error),
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import type { EmbeddingPurpose } from "@roo-code/types"
|
||||
import { OpenAICompatibleEmbedder } from "./openai-compatible"
|
||||
import { IEmbedder, EmbeddingResponse, EmbedderInfo } from "../interfaces/embedder"
|
||||
import { MAX_ITEM_TOKENS } from "../constants"
|
||||
|
|
@ -46,11 +47,11 @@ export class MistralEmbedder implements IEmbedder {
|
|||
* @param model Optional model identifier (uses constructor model if not provided)
|
||||
* @returns Promise resolving to embedding response
|
||||
*/
|
||||
async createEmbeddings(texts: string[], model?: string): Promise<EmbeddingResponse> {
|
||||
async createEmbeddings(texts: string[], model?: string, purpose?: EmbeddingPurpose): Promise<EmbeddingResponse> {
|
||||
try {
|
||||
// Use the provided model or fall back to the instance's model
|
||||
const modelToUse = model || this.modelId
|
||||
return await this.openAICompatibleEmbedder.createEmbeddings(texts, modelToUse)
|
||||
return await this.openAICompatibleEmbedder.createEmbeddings(texts, modelToUse, purpose)
|
||||
} catch (error) {
|
||||
TelemetryService.instance.captureEvent(TelemetryEventName.CODE_INDEX_ERROR, {
|
||||
error: error instanceof Error ? error.message : String(error),
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
import type { EmbeddingPurpose } from "@roo-code/types"
|
||||
import { ApiHandlerOptions } from "../../../shared/api"
|
||||
import { EmbedderInfo, EmbeddingResponse, IEmbedder } from "../interfaces"
|
||||
import { getModelQueryPrefix } from "../../../shared/embeddingModels"
|
||||
import { getModelPrefixForPurpose } from "../../../shared/embeddingModels"
|
||||
import { MAX_ITEM_TOKENS } from "../constants"
|
||||
import { t } from "../../../i18n"
|
||||
import { withValidationErrorHandling, sanitizeErrorMessage } from "../shared/validation-helpers"
|
||||
|
|
@ -35,12 +36,12 @@ export class CodeIndexOllamaEmbedder implements IEmbedder {
|
|||
* @param model - Optional model ID to override the default.
|
||||
* @returns A promise that resolves to an EmbeddingResponse containing the embeddings and usage data.
|
||||
*/
|
||||
async createEmbeddings(texts: string[], model?: string): Promise<EmbeddingResponse> {
|
||||
async createEmbeddings(texts: string[], model?: string, purpose?: EmbeddingPurpose): Promise<EmbeddingResponse> {
|
||||
const modelToUse = model || this.defaultModelId
|
||||
const url = `${this.baseUrl}/api/embed` // Endpoint as specified
|
||||
|
||||
// Apply model-specific query prefix if required
|
||||
const queryPrefix = getModelQueryPrefix("ollama", modelToUse)
|
||||
// Apply model-specific prefix based on purpose (query vs index)
|
||||
const queryPrefix = getModelPrefixForPurpose("ollama", modelToUse, purpose)
|
||||
const processedTexts = queryPrefix
|
||||
? texts.map((text, index) => {
|
||||
// Prevent double-prefixing
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import type { EmbeddingPurpose } from "@roo-code/types"
|
||||
import { OpenAI } from "openai"
|
||||
import { IEmbedder, EmbeddingResponse, EmbedderInfo } from "../interfaces/embedder"
|
||||
import {
|
||||
|
|
@ -6,7 +7,7 @@ import {
|
|||
MAX_BATCH_RETRIES as MAX_RETRIES,
|
||||
INITIAL_RETRY_DELAY_MS as INITIAL_DELAY_MS,
|
||||
} from "../constants"
|
||||
import { getDefaultModelId, getModelQueryPrefix } from "../../../shared/embeddingModels"
|
||||
import { getDefaultModelId, getModelPrefixForPurpose } from "../../../shared/embeddingModels"
|
||||
import { t } from "../../../i18n"
|
||||
import { withValidationErrorHandling, HttpError, formatEmbeddingError } from "../shared/validation-helpers"
|
||||
import { TelemetryEventName } from "@roo-code/types"
|
||||
|
|
@ -91,11 +92,11 @@ export class OpenAICompatibleEmbedder implements IEmbedder {
|
|||
* @param model Optional model identifier
|
||||
* @returns Promise resolving to embedding response
|
||||
*/
|
||||
async createEmbeddings(texts: string[], model?: string): Promise<EmbeddingResponse> {
|
||||
async createEmbeddings(texts: string[], model?: string, purpose?: EmbeddingPurpose): Promise<EmbeddingResponse> {
|
||||
const modelToUse = model || this.defaultModelId
|
||||
|
||||
// Apply model-specific query prefix if required
|
||||
const queryPrefix = getModelQueryPrefix("openai-compatible", modelToUse)
|
||||
// Apply model-specific prefix based on purpose (query vs index)
|
||||
const queryPrefix = getModelPrefixForPurpose("openai-compatible", modelToUse, purpose)
|
||||
const processedTexts = queryPrefix
|
||||
? texts.map((text, index) => {
|
||||
// Prevent double-prefixing
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import type { EmbeddingPurpose } from "@roo-code/types"
|
||||
import { OpenAI } from "openai"
|
||||
import { OpenAiNativeHandler } from "../../../api/providers/openai-native"
|
||||
import { ApiHandlerOptions } from "../../../shared/api"
|
||||
|
|
@ -8,7 +9,7 @@ import {
|
|||
MAX_BATCH_RETRIES as MAX_RETRIES,
|
||||
INITIAL_RETRY_DELAY_MS as INITIAL_DELAY_MS,
|
||||
} from "../constants"
|
||||
import { getModelQueryPrefix } from "../../../shared/embeddingModels"
|
||||
import { getModelPrefixForPurpose } from "../../../shared/embeddingModels"
|
||||
import { t } from "../../../i18n"
|
||||
import { withValidationErrorHandling, formatEmbeddingError, HttpError } from "../shared/validation-helpers"
|
||||
import { TelemetryEventName } from "@roo-code/types"
|
||||
|
|
@ -47,11 +48,11 @@ export class OpenAiEmbedder extends OpenAiNativeHandler implements IEmbedder {
|
|||
* @param model Optional model identifier
|
||||
* @returns Promise resolving to embedding response
|
||||
*/
|
||||
async createEmbeddings(texts: string[], model?: string): Promise<EmbeddingResponse> {
|
||||
async createEmbeddings(texts: string[], model?: string, purpose?: EmbeddingPurpose): Promise<EmbeddingResponse> {
|
||||
const modelToUse = model || this.defaultModelId
|
||||
|
||||
// Apply model-specific query prefix if required
|
||||
const queryPrefix = getModelQueryPrefix("openai", modelToUse)
|
||||
// Apply model-specific prefix based on purpose (query vs index)
|
||||
const queryPrefix = getModelPrefixForPurpose("openai", modelToUse, purpose)
|
||||
const processedTexts = queryPrefix
|
||||
? texts.map((text, index) => {
|
||||
// Prevent double-prefixing
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import type { EmbeddingPurpose } from "@roo-code/types"
|
||||
import { OpenAI } from "openai"
|
||||
import { IEmbedder, EmbeddingResponse, EmbedderInfo } from "../interfaces/embedder"
|
||||
import {
|
||||
|
|
@ -6,7 +7,7 @@ import {
|
|||
MAX_BATCH_RETRIES as MAX_RETRIES,
|
||||
INITIAL_RETRY_DELAY_MS as INITIAL_DELAY_MS,
|
||||
} from "../constants"
|
||||
import { getDefaultModelId, getModelQueryPrefix } from "../../../shared/embeddingModels"
|
||||
import { getDefaultModelId, getModelPrefixForPurpose } from "../../../shared/embeddingModels"
|
||||
import { t } from "../../../i18n"
|
||||
import { withValidationErrorHandling, HttpError, formatEmbeddingError } from "../shared/validation-helpers"
|
||||
import { TelemetryEventName } from "@roo-code/types"
|
||||
|
|
@ -95,11 +96,11 @@ export class OpenRouterEmbedder implements IEmbedder {
|
|||
* @param model Optional model identifier
|
||||
* @returns Promise resolving to embedding response
|
||||
*/
|
||||
async createEmbeddings(texts: string[], model?: string): Promise<EmbeddingResponse> {
|
||||
async createEmbeddings(texts: string[], model?: string, purpose?: EmbeddingPurpose): Promise<EmbeddingResponse> {
|
||||
const modelToUse = model || this.defaultModelId
|
||||
|
||||
// Apply model-specific query prefix if required
|
||||
const queryPrefix = getModelQueryPrefix("openrouter", modelToUse)
|
||||
// Apply model-specific prefix based on purpose (query vs index)
|
||||
const queryPrefix = getModelPrefixForPurpose("openrouter", modelToUse, purpose)
|
||||
const processedTexts = queryPrefix
|
||||
? texts.map((text, index) => {
|
||||
// Prevent double-prefixing
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import type { EmbeddingPurpose } from "@roo-code/types"
|
||||
import { OpenAICompatibleEmbedder } from "./openai-compatible"
|
||||
import { IEmbedder, EmbeddingResponse, EmbedderInfo } from "../interfaces/embedder"
|
||||
import { MAX_ITEM_TOKENS } from "../constants"
|
||||
|
|
@ -55,11 +56,11 @@ export class VercelAiGatewayEmbedder implements IEmbedder {
|
|||
* @param model Optional model identifier (uses constructor model if not provided)
|
||||
* @returns Promise resolving to embedding response
|
||||
*/
|
||||
async createEmbeddings(texts: string[], model?: string): Promise<EmbeddingResponse> {
|
||||
async createEmbeddings(texts: string[], model?: string, purpose?: EmbeddingPurpose): Promise<EmbeddingResponse> {
|
||||
try {
|
||||
// Use the provided model or fall back to the instance's model
|
||||
const modelToUse = model || this.modelId
|
||||
return await this.openAICompatibleEmbedder.createEmbeddings(texts, modelToUse)
|
||||
return await this.openAICompatibleEmbedder.createEmbeddings(texts, modelToUse, purpose)
|
||||
} catch (error) {
|
||||
TelemetryService.instance.captureEvent(TelemetryEventName.CODE_INDEX_ERROR, {
|
||||
error: error instanceof Error ? error.message : String(error),
|
||||
|
|
|
|||
|
|
@ -1,3 +1,5 @@
|
|||
import type { EmbeddingPurpose } from "@roo-code/types"
|
||||
|
||||
/**
|
||||
* Interface for code index embedders.
|
||||
* This interface is implemented by both OpenAI and Ollama embedders.
|
||||
|
|
@ -7,9 +9,12 @@ export interface IEmbedder {
|
|||
* Creates embeddings for the given texts.
|
||||
* @param texts Array of text strings to create embeddings for
|
||||
* @param model Optional model ID to use for embeddings
|
||||
* @param purpose Optional embedding purpose ("index" for document indexing, "query" for search queries).
|
||||
* When "index", the model's documentPrefix is applied (if any).
|
||||
* When "query" or undefined, the model's queryPrefix is applied (if any) for backward compatibility.
|
||||
* @returns Promise resolving to an EmbeddingResponse
|
||||
*/
|
||||
createEmbeddings(texts: string[], model?: string): Promise<EmbeddingResponse>
|
||||
createEmbeddings(texts: string[], model?: string, purpose?: EmbeddingPurpose): Promise<EmbeddingResponse>
|
||||
|
||||
/**
|
||||
* Validates the embedder configuration by testing connectivity and credentials.
|
||||
|
|
|
|||
|
|
@ -566,7 +566,7 @@ export class FileWatcher implements IFileWatcher {
|
|||
let pointsToUpsert: PointStruct[] = []
|
||||
if (this.embedder && blocks.length > 0) {
|
||||
const texts = blocks.map((block) => block.content)
|
||||
const { embeddings } = await this.embedder.createEmbeddings(texts)
|
||||
const { embeddings } = await this.embedder.createEmbeddings(texts, undefined, "index")
|
||||
|
||||
pointsToUpsert = blocks.map((block, index) => {
|
||||
const normalizedAbsolutePath = generateNormalizedAbsolutePath(block.file_path, this.workspacePath)
|
||||
|
|
|
|||
|
|
@ -442,8 +442,8 @@ export class DirectoryScanner implements IDirectoryScanner {
|
|||
}
|
||||
// --- End Deletion Step ---
|
||||
|
||||
// Create embeddings for batch
|
||||
const { embeddings } = await this.embedder.createEmbeddings(batchTexts)
|
||||
// Create embeddings for batch (purpose: "index" for document indexing)
|
||||
const { embeddings } = await this.embedder.createEmbeddings(batchTexts, undefined, "index")
|
||||
|
||||
// Prepare points for Qdrant
|
||||
const points = batchBlocks.map((block, index) => {
|
||||
|
|
|
|||
|
|
@ -41,8 +41,8 @@ export class CodeIndexSearchService {
|
|||
}
|
||||
|
||||
try {
|
||||
// Generate embedding for query
|
||||
const embeddingResponse = await this.embedder.createEmbeddings([query])
|
||||
// Generate embedding for query (purpose: "query" for search queries)
|
||||
const embeddingResponse = await this.embedder.createEmbeddings([query], undefined, "query")
|
||||
const vector = embeddingResponse?.embeddings[0]
|
||||
if (!vector) {
|
||||
throw new Error("Failed to generate embedding for query.")
|
||||
|
|
|
|||
|
|
@ -2,6 +2,9 @@ import { describe, it, expect } from "vitest"
|
|||
import {
|
||||
getModelDimension,
|
||||
getModelScoreThreshold,
|
||||
getModelQueryPrefix,
|
||||
getModelDocumentPrefix,
|
||||
getModelPrefixForPurpose,
|
||||
getDefaultModelId,
|
||||
EMBEDDING_MODEL_PROFILES,
|
||||
} from "../embeddingModels"
|
||||
|
|
@ -92,4 +95,75 @@ describe("embeddingModels", () => {
|
|||
expect(defaultModel).toBe("codestral-embed-2505")
|
||||
})
|
||||
})
|
||||
|
||||
describe("getModelQueryPrefix", () => {
|
||||
it("should return queryPrefix for nomic-embed-code on ollama", () => {
|
||||
const prefix = getModelQueryPrefix("ollama", "nomic-embed-code")
|
||||
expect(prefix).toBe("Represent this query for searching relevant code: ")
|
||||
})
|
||||
|
||||
it("should return queryPrefix for nomic-embed-code on openai-compatible", () => {
|
||||
const prefix = getModelQueryPrefix("openai-compatible", "nomic-embed-code")
|
||||
expect(prefix).toBe("Represent this query for searching relevant code: ")
|
||||
})
|
||||
|
||||
it("should return undefined for models without queryPrefix", () => {
|
||||
const prefix = getModelQueryPrefix("openai", "text-embedding-3-small")
|
||||
expect(prefix).toBeUndefined()
|
||||
})
|
||||
|
||||
it("should return undefined for unknown model", () => {
|
||||
const prefix = getModelQueryPrefix("ollama", "unknown-model")
|
||||
expect(prefix).toBeUndefined()
|
||||
})
|
||||
|
||||
it("should return undefined for unknown provider", () => {
|
||||
const prefix = getModelQueryPrefix("unknown-provider" as any, "some-model")
|
||||
expect(prefix).toBeUndefined()
|
||||
})
|
||||
})
|
||||
|
||||
describe("getModelDocumentPrefix", () => {
|
||||
it("should return undefined for nomic-embed-code on ollama (no prefix for indexing)", () => {
|
||||
const prefix = getModelDocumentPrefix("ollama", "nomic-embed-code")
|
||||
expect(prefix).toBeUndefined()
|
||||
})
|
||||
|
||||
it("should return undefined for nomic-embed-code on openai-compatible (no prefix for indexing)", () => {
|
||||
const prefix = getModelDocumentPrefix("openai-compatible", "nomic-embed-code")
|
||||
expect(prefix).toBeUndefined()
|
||||
})
|
||||
|
||||
it("should return undefined for models without documentPrefix", () => {
|
||||
const prefix = getModelDocumentPrefix("openai", "text-embedding-3-small")
|
||||
expect(prefix).toBeUndefined()
|
||||
})
|
||||
|
||||
it("should return undefined for unknown provider", () => {
|
||||
const prefix = getModelDocumentPrefix("unknown-provider" as any, "some-model")
|
||||
expect(prefix).toBeUndefined()
|
||||
})
|
||||
})
|
||||
|
||||
describe("getModelPrefixForPurpose", () => {
|
||||
it("should return queryPrefix when purpose is 'query'", () => {
|
||||
const prefix = getModelPrefixForPurpose("ollama", "nomic-embed-code", "query")
|
||||
expect(prefix).toBe("Represent this query for searching relevant code: ")
|
||||
})
|
||||
|
||||
it("should return documentPrefix (undefined) when purpose is 'index'", () => {
|
||||
const prefix = getModelPrefixForPurpose("ollama", "nomic-embed-code", "index")
|
||||
expect(prefix).toBeUndefined()
|
||||
})
|
||||
|
||||
it("should fall back to queryPrefix when purpose is undefined (backward compatibility)", () => {
|
||||
const prefix = getModelPrefixForPurpose("ollama", "nomic-embed-code", undefined)
|
||||
expect(prefix).toBe("Represent this query for searching relevant code: ")
|
||||
})
|
||||
|
||||
it("should return undefined for models without any prefix regardless of purpose", () => {
|
||||
expect(getModelPrefixForPurpose("openai", "text-embedding-3-small", "query")).toBeUndefined()
|
||||
expect(getModelPrefixForPurpose("openai", "text-embedding-3-small", "index")).toBeUndefined()
|
||||
})
|
||||
})
|
||||
})
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@
|
|||
* Defines profiles for different embedding models, including their dimensions.
|
||||
*/
|
||||
|
||||
import type { EmbedderProvider, EmbeddingModelProfiles } from "@roo-code/types"
|
||||
import type { EmbedderProvider, EmbeddingModelProfiles, EmbeddingPurpose } from "@roo-code/types"
|
||||
|
||||
// Example profiles - expand this list as needed
|
||||
export const EMBEDDING_MODEL_PROFILES: EmbeddingModelProfiles = {
|
||||
|
|
@ -17,6 +17,8 @@ export const EMBEDDING_MODEL_PROFILES: EmbeddingModelProfiles = {
|
|||
dimension: 3584,
|
||||
scoreThreshold: 0.15,
|
||||
queryPrefix: "Represent this query for searching relevant code: ",
|
||||
// documentPrefix is intentionally omitted (undefined) -- nomic-embed-code
|
||||
// does not require a prefix when embedding documents for indexing.
|
||||
},
|
||||
"mxbai-embed-large": { dimension: 1024, scoreThreshold: 0.4 },
|
||||
"all-minilm": { dimension: 384, scoreThreshold: 0.4 },
|
||||
|
|
@ -31,6 +33,8 @@ export const EMBEDDING_MODEL_PROFILES: EmbeddingModelProfiles = {
|
|||
dimension: 3584,
|
||||
scoreThreshold: 0.15,
|
||||
queryPrefix: "Represent this query for searching relevant code: ",
|
||||
// documentPrefix is intentionally omitted (undefined) -- nomic-embed-code
|
||||
// does not require a prefix when embedding documents for indexing.
|
||||
},
|
||||
},
|
||||
gemini: {
|
||||
|
|
@ -144,6 +148,45 @@ export function getModelQueryPrefix(provider: EmbedderProvider, modelId: string)
|
|||
return modelProfile?.queryPrefix
|
||||
}
|
||||
|
||||
/**
|
||||
* Retrieves the document prefix for a given provider and model ID.
|
||||
* This prefix is used when embedding documents for indexing (as opposed to search queries).
|
||||
* @param provider The embedder provider (e.g., "ollama").
|
||||
* @param modelId The specific model ID (e.g., "nomic-embed-code").
|
||||
* @returns The document prefix or undefined if the model doesn't require one.
|
||||
*/
|
||||
export function getModelDocumentPrefix(provider: EmbedderProvider, modelId: string): string | undefined {
|
||||
const providerProfiles = EMBEDDING_MODEL_PROFILES[provider]
|
||||
if (!providerProfiles) {
|
||||
return undefined
|
||||
}
|
||||
|
||||
const modelProfile = providerProfiles[modelId]
|
||||
return modelProfile?.documentPrefix
|
||||
}
|
||||
|
||||
/**
|
||||
* Retrieves the appropriate prefix for a given purpose (indexing or querying).
|
||||
* When purpose is "index", returns the documentPrefix.
|
||||
* When purpose is "query", returns the queryPrefix.
|
||||
* When purpose is undefined, falls back to queryPrefix for backward compatibility.
|
||||
* @param provider The embedder provider.
|
||||
* @param modelId The specific model ID.
|
||||
* @param purpose The embedding purpose ("index" or "query").
|
||||
* @returns The appropriate prefix or undefined.
|
||||
*/
|
||||
export function getModelPrefixForPurpose(
|
||||
provider: EmbedderProvider,
|
||||
modelId: string,
|
||||
purpose?: EmbeddingPurpose,
|
||||
): string | undefined {
|
||||
if (purpose === "index") {
|
||||
return getModelDocumentPrefix(provider, modelId)
|
||||
}
|
||||
// For "query" or undefined (backward compatibility), use queryPrefix
|
||||
return getModelQueryPrefix(provider, modelId)
|
||||
}
|
||||
|
||||
/**
|
||||
* Gets the default *specific* embedding model ID based on the provider.
|
||||
* Does not include the provider prefix.
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue