mirror of
https://github.com/RooVetGit/Roo-Code.git
synced 2026-08-28 05:27:24 +00:00
feat: add custom headers support for OpenAI compatible embeddings
- Add codebaseIndexOpenAiCompatibleHeaders field to types - Update config manager to handle custom headers - Modify OpenAICompatibleEmbedder to accept and use custom headers - Update service factory to pass headers to embedder - Add comprehensive test coverage for custom headers Fixes #8909
This commit is contained in:
parent
ff0c65af10
commit
d089fb28ae
6 changed files with 142 additions and 13 deletions
|
|
@ -36,6 +36,7 @@ export const codebaseIndexConfigSchema = z.object({
|
|||
// OpenAI Compatible specific fields
|
||||
codebaseIndexOpenAiCompatibleBaseUrl: z.string().optional(),
|
||||
codebaseIndexOpenAiCompatibleModelDimension: z.number().optional(),
|
||||
codebaseIndexOpenAiCompatibleHeaders: z.record(z.string()).optional(),
|
||||
})
|
||||
|
||||
export type CodebaseIndexConfig = z.infer<typeof codebaseIndexConfigSchema>
|
||||
|
|
@ -65,6 +66,7 @@ export const codebaseIndexProviderSchema = z.object({
|
|||
codebaseIndexOpenAiCompatibleBaseUrl: z.string().optional(),
|
||||
codebaseIndexOpenAiCompatibleApiKey: z.string().optional(),
|
||||
codebaseIndexOpenAiCompatibleModelDimension: z.number().optional(),
|
||||
codebaseIndexOpenAiCompatibleHeaders: z.record(z.string()).optional(),
|
||||
codebaseIndexGeminiApiKey: z.string().optional(),
|
||||
codebaseIndexMistralApiKey: z.string().optional(),
|
||||
codebaseIndexVercelAiGatewayApiKey: z.string().optional(),
|
||||
|
|
|
|||
|
|
@ -16,7 +16,7 @@ export class CodeIndexConfigManager {
|
|||
private modelDimension?: number
|
||||
private openAiOptions?: ApiHandlerOptions
|
||||
private ollamaOptions?: ApiHandlerOptions
|
||||
private openAiCompatibleOptions?: { baseUrl: string; apiKey: string }
|
||||
private openAiCompatibleOptions?: { baseUrl: string; apiKey: string; headers?: Record<string, string> }
|
||||
private geminiOptions?: { apiKey: string }
|
||||
private mistralOptions?: { apiKey: string }
|
||||
private vercelAiGatewayOptions?: { apiKey: string }
|
||||
|
|
@ -68,6 +68,7 @@ export class CodeIndexConfigManager {
|
|||
// Fix: Read OpenAI Compatible settings from the correct location within codebaseIndexConfig
|
||||
const openAiCompatibleBaseUrl = codebaseIndexConfig.codebaseIndexOpenAiCompatibleBaseUrl ?? ""
|
||||
const openAiCompatibleApiKey = this.contextProxy?.getSecret("codebaseIndexOpenAiCompatibleApiKey") ?? ""
|
||||
const openAiCompatibleHeaders = codebaseIndexConfig.codebaseIndexOpenAiCompatibleHeaders ?? undefined
|
||||
const geminiApiKey = this.contextProxy?.getSecret("codebaseIndexGeminiApiKey") ?? ""
|
||||
const mistralApiKey = this.contextProxy?.getSecret("codebaseIndexMistralApiKey") ?? ""
|
||||
const vercelAiGatewayApiKey = this.contextProxy?.getSecret("codebaseIndexVercelAiGatewayApiKey") ?? ""
|
||||
|
|
@ -123,6 +124,7 @@ export class CodeIndexConfigManager {
|
|||
? {
|
||||
baseUrl: openAiCompatibleBaseUrl,
|
||||
apiKey: openAiCompatibleApiKey,
|
||||
headers: openAiCompatibleHeaders,
|
||||
}
|
||||
: undefined
|
||||
|
||||
|
|
@ -143,7 +145,7 @@ export class CodeIndexConfigManager {
|
|||
modelDimension?: number
|
||||
openAiOptions?: ApiHandlerOptions
|
||||
ollamaOptions?: ApiHandlerOptions
|
||||
openAiCompatibleOptions?: { baseUrl: string; apiKey: string }
|
||||
openAiCompatibleOptions?: { baseUrl: string; apiKey: string; headers?: Record<string, string> }
|
||||
geminiOptions?: { apiKey: string }
|
||||
mistralOptions?: { apiKey: string }
|
||||
vercelAiGatewayOptions?: { apiKey: string }
|
||||
|
|
|
|||
|
|
@ -112,6 +112,31 @@ describe("OpenAICompatibleEmbedder", () => {
|
|||
expect(embedder).toBeDefined()
|
||||
})
|
||||
|
||||
it("should create embedder with custom headers", () => {
|
||||
const customHeaders = {
|
||||
"X-Custom-Header": "custom-value",
|
||||
"X-Another-Header": "another-value",
|
||||
}
|
||||
embedder = new OpenAICompatibleEmbedder(testBaseUrl, testApiKey, testModelId, undefined, customHeaders)
|
||||
|
||||
expect(MockedOpenAI).toHaveBeenCalledWith({
|
||||
baseURL: testBaseUrl,
|
||||
apiKey: testApiKey,
|
||||
defaultHeaders: customHeaders,
|
||||
})
|
||||
expect(embedder).toBeDefined()
|
||||
})
|
||||
|
||||
it("should create embedder without custom headers when not provided", () => {
|
||||
embedder = new OpenAICompatibleEmbedder(testBaseUrl, testApiKey, testModelId, undefined, undefined)
|
||||
|
||||
expect(MockedOpenAI).toHaveBeenCalledWith({
|
||||
baseURL: testBaseUrl,
|
||||
apiKey: testApiKey,
|
||||
})
|
||||
expect(embedder).toBeDefined()
|
||||
})
|
||||
|
||||
it("should throw error when baseUrl is missing", () => {
|
||||
expect(() => new OpenAICompatibleEmbedder("", testApiKey, testModelId)).toThrow(
|
||||
"embeddings:validation.baseUrlRequired",
|
||||
|
|
@ -813,6 +838,81 @@ describe("OpenAICompatibleEmbedder", () => {
|
|||
expect(baseResult.embeddings[0]).toEqual([0.4, 0.5, 0.6])
|
||||
})
|
||||
|
||||
it("should include custom headers in direct fetch requests", async () => {
|
||||
const testTexts = ["Test text"]
|
||||
const customHeaders = {
|
||||
"X-Custom-Header": "custom-value",
|
||||
"X-API-Version": "v2",
|
||||
}
|
||||
const base64String = createBase64Embedding([0.1, 0.2, 0.3])
|
||||
|
||||
// Test Azure URL with custom headers (direct fetch)
|
||||
const azureEmbedder = new OpenAICompatibleEmbedder(
|
||||
azureUrl,
|
||||
testApiKey,
|
||||
testModelId,
|
||||
undefined,
|
||||
customHeaders,
|
||||
)
|
||||
const mockFetchResponse = createMockResponse({
|
||||
data: [{ embedding: base64String }],
|
||||
usage: { prompt_tokens: 10, total_tokens: 15 },
|
||||
})
|
||||
;(global.fetch as MockedFunction<typeof fetch>).mockResolvedValue(mockFetchResponse as any)
|
||||
|
||||
const azureResult = await azureEmbedder.createEmbeddings(testTexts)
|
||||
expect(global.fetch).toHaveBeenCalledWith(
|
||||
azureUrl,
|
||||
expect.objectContaining({
|
||||
method: "POST",
|
||||
headers: expect.objectContaining({
|
||||
"Content-Type": "application/json",
|
||||
"api-key": testApiKey,
|
||||
Authorization: `Bearer ${testApiKey}`,
|
||||
"X-Custom-Header": "custom-value",
|
||||
"X-API-Version": "v2",
|
||||
}),
|
||||
}),
|
||||
)
|
||||
expect(mockEmbeddingsCreate).not.toHaveBeenCalled()
|
||||
expectEmbeddingValues(azureResult.embeddings[0], [0.1, 0.2, 0.3])
|
||||
})
|
||||
|
||||
it("should handle custom headers that override default headers", async () => {
|
||||
const testTexts = ["Test text"]
|
||||
const customHeaders = {
|
||||
"api-key": "override-key", // Override the default api-key
|
||||
"X-Custom-Header": "custom-value",
|
||||
}
|
||||
const base64String = createBase64Embedding([0.1, 0.2, 0.3])
|
||||
|
||||
const azureEmbedder = new OpenAICompatibleEmbedder(
|
||||
azureUrl,
|
||||
testApiKey,
|
||||
testModelId,
|
||||
undefined,
|
||||
customHeaders,
|
||||
)
|
||||
const mockFetchResponse = createMockResponse({
|
||||
data: [{ embedding: base64String }],
|
||||
usage: { prompt_tokens: 10, total_tokens: 15 },
|
||||
})
|
||||
;(global.fetch as MockedFunction<typeof fetch>).mockResolvedValue(mockFetchResponse as any)
|
||||
|
||||
const azureResult = await azureEmbedder.createEmbeddings(testTexts)
|
||||
expect(global.fetch).toHaveBeenCalledWith(
|
||||
azureUrl,
|
||||
expect.objectContaining({
|
||||
method: "POST",
|
||||
headers: expect.objectContaining({
|
||||
"api-key": "override-key", // Custom header overrides default
|
||||
"X-Custom-Header": "custom-value",
|
||||
}),
|
||||
}),
|
||||
)
|
||||
expectEmbeddingValues(azureResult.embeddings[0], [0.1, 0.2, 0.3])
|
||||
})
|
||||
|
||||
it.each([
|
||||
[401, "Authentication failed. Please check your API key."],
|
||||
[500, "Failed to create embeddings after 3 attempts"],
|
||||
|
|
|
|||
|
|
@ -37,6 +37,7 @@ export class OpenAICompatibleEmbedder implements IEmbedder {
|
|||
private readonly defaultModelId: string
|
||||
private readonly baseUrl: string
|
||||
private readonly apiKey: string
|
||||
private readonly customHeaders?: Record<string, string>
|
||||
private readonly isFullUrl: boolean
|
||||
private readonly maxItemTokens: number
|
||||
|
||||
|
|
@ -56,8 +57,15 @@ export class OpenAICompatibleEmbedder implements IEmbedder {
|
|||
* @param apiKey The API key for authentication
|
||||
* @param modelId Optional model identifier (defaults to "text-embedding-3-small")
|
||||
* @param maxItemTokens Optional maximum tokens per item (defaults to MAX_ITEM_TOKENS)
|
||||
* @param customHeaders Optional custom headers to include in requests
|
||||
*/
|
||||
constructor(baseUrl: string, apiKey: string, modelId?: string, maxItemTokens?: number) {
|
||||
constructor(
|
||||
baseUrl: string,
|
||||
apiKey: string,
|
||||
modelId?: string,
|
||||
maxItemTokens?: number,
|
||||
customHeaders?: Record<string, string>,
|
||||
) {
|
||||
if (!baseUrl) {
|
||||
throw new Error(t("embeddings:validation.baseUrlRequired"))
|
||||
}
|
||||
|
|
@ -67,13 +75,21 @@ export class OpenAICompatibleEmbedder implements IEmbedder {
|
|||
|
||||
this.baseUrl = baseUrl
|
||||
this.apiKey = apiKey
|
||||
this.customHeaders = customHeaders
|
||||
|
||||
// Wrap OpenAI client creation to handle invalid API key characters
|
||||
try {
|
||||
this.embeddingsClient = new OpenAI({
|
||||
// If custom headers are provided, we need to use defaultHeaders in OpenAI config
|
||||
const openAIConfig: any = {
|
||||
baseURL: baseUrl,
|
||||
apiKey: apiKey,
|
||||
})
|
||||
}
|
||||
|
||||
if (customHeaders) {
|
||||
openAIConfig.defaultHeaders = customHeaders
|
||||
}
|
||||
|
||||
this.embeddingsClient = new OpenAI(openAIConfig)
|
||||
} catch (error) {
|
||||
// Use the error handler to transform ByteString conversion errors
|
||||
throw handleOpenAIError(error, "OpenAI Compatible")
|
||||
|
|
@ -204,15 +220,22 @@ export class OpenAICompatibleEmbedder implements IEmbedder {
|
|||
batchTexts: string[],
|
||||
model: string,
|
||||
): Promise<OpenAIEmbeddingResponse> {
|
||||
const headers: Record<string, string> = {
|
||||
"Content-Type": "application/json",
|
||||
// Azure OpenAI uses 'api-key' header, while OpenAI uses 'Authorization'
|
||||
// We'll try 'api-key' first for Azure compatibility
|
||||
"api-key": this.apiKey,
|
||||
Authorization: `Bearer ${this.apiKey}`,
|
||||
}
|
||||
|
||||
// Add custom headers if provided
|
||||
if (this.customHeaders) {
|
||||
Object.assign(headers, this.customHeaders)
|
||||
}
|
||||
|
||||
const response = await fetch(url, {
|
||||
method: "POST",
|
||||
headers: {
|
||||
"Content-Type": "application/json",
|
||||
// Azure OpenAI uses 'api-key' header, while OpenAI uses 'Authorization'
|
||||
// We'll try 'api-key' first for Azure compatibility
|
||||
"api-key": this.apiKey,
|
||||
Authorization: `Bearer ${this.apiKey}`,
|
||||
},
|
||||
headers,
|
||||
body: JSON.stringify({
|
||||
input: batchTexts,
|
||||
model: model,
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ export interface CodeIndexConfig {
|
|||
modelDimension?: number // Generic dimension property for all providers
|
||||
openAiOptions?: ApiHandlerOptions
|
||||
ollamaOptions?: ApiHandlerOptions
|
||||
openAiCompatibleOptions?: { baseUrl: string; apiKey: string }
|
||||
openAiCompatibleOptions?: { baseUrl: string; apiKey: string; headers?: Record<string, string> }
|
||||
geminiOptions?: { apiKey: string }
|
||||
mistralOptions?: { apiKey: string }
|
||||
vercelAiGatewayOptions?: { apiKey: string }
|
||||
|
|
|
|||
|
|
@ -63,6 +63,8 @@ export class CodeIndexServiceFactory {
|
|||
config.openAiCompatibleOptions.baseUrl,
|
||||
config.openAiCompatibleOptions.apiKey,
|
||||
config.modelId,
|
||||
undefined, // maxItemTokens (use default)
|
||||
config.openAiCompatibleOptions.headers,
|
||||
)
|
||||
} else if (provider === "gemini") {
|
||||
if (!config.geminiOptions?.apiKey) {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue