mirror of
https://github.com/RooVetGit/Roo-Code.git
synced 2026-08-28 05:27:24 +00:00
feat: add per-profile GOOGLE_CLOUD_PROJECT support for Gemini and Vertex providers
- Add googleCloudProject field to provider settings schema for gemini, vertex, and gemini-cli providers - Update GeminiHandler and AnthropicVertexHandler to set GOOGLE_CLOUD_PROJECT environment variable during API calls - Add utility functions for managing Google Cloud environment variables - Include comprehensive tests for environment variable handling Fixes #5799
This commit is contained in:
parent
6cf376f832
commit
029df44707
6 changed files with 659 additions and 237 deletions
|
|
@ -122,6 +122,7 @@ const vertexSchema = apiModelIdProviderModelSchema.extend({
|
|||
vertexJsonCredentials: z.string().optional(),
|
||||
vertexProjectId: z.string().optional(),
|
||||
vertexRegion: z.string().optional(),
|
||||
googleCloudProject: z.string().optional(),
|
||||
})
|
||||
|
||||
const openAiSchema = baseProviderSettingsSchema.extend({
|
||||
|
|
@ -164,11 +165,13 @@ const lmStudioSchema = baseProviderSettingsSchema.extend({
|
|||
const geminiSchema = apiModelIdProviderModelSchema.extend({
|
||||
geminiApiKey: z.string().optional(),
|
||||
googleGeminiBaseUrl: z.string().optional(),
|
||||
googleCloudProject: z.string().optional(),
|
||||
})
|
||||
|
||||
const geminiCliSchema = apiModelIdProviderModelSchema.extend({
|
||||
geminiCliOAuthPath: z.string().optional(),
|
||||
geminiCliProjectId: z.string().optional(),
|
||||
googleCloudProject: z.string().optional(),
|
||||
})
|
||||
|
||||
const openAiNativeSchema = apiModelIdProviderModelSchema.extend({
|
||||
|
|
|
|||
169
src/api/providers/__tests__/gemini-google-cloud-project.spec.ts
Normal file
169
src/api/providers/__tests__/gemini-google-cloud-project.spec.ts
Normal file
|
|
@ -0,0 +1,169 @@
|
|||
// npx vitest run src/api/providers/__tests__/gemini-google-cloud-project.spec.ts
|
||||
|
||||
import { describe, it, expect, beforeEach, afterEach, vi } from "vitest"
|
||||
|
||||
import { GeminiHandler } from "../gemini"
|
||||
|
||||
// Mock the @google/genai module
|
||||
const mockGenerateContentStream = vi.fn()
|
||||
const mockGenerateContent = vi.fn()
|
||||
const mockCountTokens = vi.fn()
|
||||
|
||||
vi.mock("@google/genai", () => ({
|
||||
GoogleGenAI: vi.fn().mockImplementation(() => ({
|
||||
models: {
|
||||
generateContentStream: mockGenerateContentStream,
|
||||
generateContent: mockGenerateContent,
|
||||
countTokens: mockCountTokens,
|
||||
},
|
||||
})),
|
||||
}))
|
||||
|
||||
describe("GeminiHandler Google Cloud Project", () => {
|
||||
let originalValue: string | undefined
|
||||
|
||||
beforeEach(() => {
|
||||
// Store the original value
|
||||
originalValue = process.env.GOOGLE_CLOUD_PROJECT
|
||||
// Clean up the environment variable
|
||||
delete process.env.GOOGLE_CLOUD_PROJECT
|
||||
})
|
||||
|
||||
afterEach(() => {
|
||||
// Restore the original value
|
||||
if (originalValue !== undefined) {
|
||||
process.env.GOOGLE_CLOUD_PROJECT = originalValue
|
||||
} else {
|
||||
delete process.env.GOOGLE_CLOUD_PROJECT
|
||||
}
|
||||
})
|
||||
|
||||
it("should set GOOGLE_CLOUD_PROJECT during createMessage when specified in profile", async () => {
|
||||
const testProjectId = "test-project-123"
|
||||
const handler = new GeminiHandler({
|
||||
geminiApiKey: "test-key",
|
||||
googleCloudProject: testProjectId,
|
||||
})
|
||||
|
||||
// Mock the generateContentStream to capture the environment variable
|
||||
let capturedEnvValue: string | undefined
|
||||
const mockStream = {
|
||||
async *[Symbol.asyncIterator]() {
|
||||
capturedEnvValue = process.env.GOOGLE_CLOUD_PROJECT
|
||||
yield { text: "test response" }
|
||||
},
|
||||
}
|
||||
|
||||
mockGenerateContentStream.mockResolvedValue(mockStream)
|
||||
|
||||
// Execute createMessage
|
||||
const generator = handler.createMessage("test system", [{ role: "user", content: "test message" }])
|
||||
|
||||
// Consume the generator
|
||||
const results = []
|
||||
for await (const chunk of generator) {
|
||||
results.push(chunk)
|
||||
}
|
||||
|
||||
// Verify the environment variable was set during execution
|
||||
expect(capturedEnvValue).toBe(testProjectId)
|
||||
// Verify it was restored after execution
|
||||
expect(process.env.GOOGLE_CLOUD_PROJECT).toBeUndefined()
|
||||
})
|
||||
|
||||
it("should set GOOGLE_CLOUD_PROJECT during completePrompt when specified in profile", async () => {
|
||||
const testProjectId = "test-project-456"
|
||||
const handler = new GeminiHandler({
|
||||
geminiApiKey: "test-key",
|
||||
googleCloudProject: testProjectId,
|
||||
})
|
||||
|
||||
// Mock the generateContent to capture the environment variable
|
||||
let capturedEnvValue: string | undefined
|
||||
mockGenerateContent.mockImplementation(async () => {
|
||||
capturedEnvValue = process.env.GOOGLE_CLOUD_PROJECT
|
||||
return { text: "test response" } as any
|
||||
})
|
||||
|
||||
// Execute completePrompt
|
||||
await handler.completePrompt("test prompt")
|
||||
|
||||
// Verify the environment variable was set during execution
|
||||
expect(capturedEnvValue).toBe(testProjectId)
|
||||
// Verify it was restored after execution
|
||||
expect(process.env.GOOGLE_CLOUD_PROJECT).toBeUndefined()
|
||||
})
|
||||
|
||||
it("should set GOOGLE_CLOUD_PROJECT during countTokens when specified in profile", async () => {
|
||||
const testProjectId = "test-project-789"
|
||||
const handler = new GeminiHandler({
|
||||
geminiApiKey: "test-key",
|
||||
googleCloudProject: testProjectId,
|
||||
})
|
||||
|
||||
// Mock the countTokens to capture the environment variable
|
||||
let capturedEnvValue: string | undefined
|
||||
mockCountTokens.mockImplementation(async () => {
|
||||
capturedEnvValue = process.env.GOOGLE_CLOUD_PROJECT
|
||||
return { totalTokens: 100 }
|
||||
})
|
||||
|
||||
// Execute countTokens
|
||||
await handler.countTokens([{ type: "text", text: "test content" }])
|
||||
|
||||
// Verify the environment variable was set during execution
|
||||
expect(capturedEnvValue).toBe(testProjectId)
|
||||
// Verify it was restored after execution
|
||||
expect(process.env.GOOGLE_CLOUD_PROJECT).toBeUndefined()
|
||||
})
|
||||
|
||||
it("should not modify GOOGLE_CLOUD_PROJECT when not specified in profile", async () => {
|
||||
const originalProjectId = "original-project"
|
||||
process.env.GOOGLE_CLOUD_PROJECT = originalProjectId
|
||||
|
||||
const handler = new GeminiHandler({
|
||||
geminiApiKey: "test-key",
|
||||
// No googleCloudProject specified
|
||||
})
|
||||
|
||||
// Mock the generateContent
|
||||
let capturedEnvValue: string | undefined
|
||||
mockGenerateContent.mockImplementation(async () => {
|
||||
capturedEnvValue = process.env.GOOGLE_CLOUD_PROJECT
|
||||
return { text: "test response" } as any
|
||||
})
|
||||
|
||||
// Execute completePrompt
|
||||
await handler.completePrompt("test prompt")
|
||||
|
||||
// Verify the environment variable was not modified
|
||||
expect(capturedEnvValue).toBe(originalProjectId)
|
||||
expect(process.env.GOOGLE_CLOUD_PROJECT).toBe(originalProjectId)
|
||||
})
|
||||
|
||||
it("should restore original GOOGLE_CLOUD_PROJECT value after execution", async () => {
|
||||
const originalProjectId = "original-project"
|
||||
const testProjectId = "test-project-123"
|
||||
process.env.GOOGLE_CLOUD_PROJECT = originalProjectId
|
||||
|
||||
const handler = new GeminiHandler({
|
||||
geminiApiKey: "test-key",
|
||||
googleCloudProject: testProjectId,
|
||||
})
|
||||
|
||||
// Mock the generateContent
|
||||
let capturedEnvValue: string | undefined
|
||||
mockGenerateContent.mockImplementation(async () => {
|
||||
capturedEnvValue = process.env.GOOGLE_CLOUD_PROJECT
|
||||
return { text: "test response" } as any
|
||||
})
|
||||
|
||||
// Execute completePrompt
|
||||
await handler.completePrompt("test prompt")
|
||||
|
||||
// Verify the test project was set during execution
|
||||
expect(capturedEnvValue).toBe(testProjectId)
|
||||
// Verify the original value was restored
|
||||
expect(process.env.GOOGLE_CLOUD_PROJECT).toBe(originalProjectId)
|
||||
})
|
||||
})
|
||||
|
|
@ -12,6 +12,7 @@ import {
|
|||
|
||||
import { ApiHandlerOptions } from "../../shared/api"
|
||||
import { safeJsonParse } from "../../shared/safeJsonParse"
|
||||
import { withGoogleCloudProject, withGoogleCloudProjectSync } from "../../utils/googleCloudEnv"
|
||||
|
||||
import { ApiStream } from "../transform/stream"
|
||||
import { addCacheBreakpoints } from "../transform/caching/vertex"
|
||||
|
|
@ -34,27 +35,30 @@ export class AnthropicVertexHandler extends BaseProvider implements SingleComple
|
|||
const projectId = this.options.vertexProjectId ?? "not-provided"
|
||||
const region = this.options.vertexRegion ?? "us-east5"
|
||||
|
||||
if (this.options.vertexJsonCredentials) {
|
||||
this.client = new AnthropicVertex({
|
||||
projectId,
|
||||
region,
|
||||
googleAuth: new GoogleAuth({
|
||||
scopes: ["https://www.googleapis.com/auth/cloud-platform"],
|
||||
credentials: safeJsonParse<JWTInput>(this.options.vertexJsonCredentials, undefined),
|
||||
}),
|
||||
})
|
||||
} else if (this.options.vertexKeyFile) {
|
||||
this.client = new AnthropicVertex({
|
||||
projectId,
|
||||
region,
|
||||
googleAuth: new GoogleAuth({
|
||||
scopes: ["https://www.googleapis.com/auth/cloud-platform"],
|
||||
keyFile: this.options.vertexKeyFile,
|
||||
}),
|
||||
})
|
||||
} else {
|
||||
this.client = new AnthropicVertex({ projectId, region })
|
||||
}
|
||||
// Set GOOGLE_CLOUD_PROJECT environment variable if specified in profile
|
||||
this.client = withGoogleCloudProjectSync(this.options.googleCloudProject, () => {
|
||||
if (this.options.vertexJsonCredentials) {
|
||||
return new AnthropicVertex({
|
||||
projectId,
|
||||
region,
|
||||
googleAuth: new GoogleAuth({
|
||||
scopes: ["https://www.googleapis.com/auth/cloud-platform"],
|
||||
credentials: safeJsonParse<JWTInput>(this.options.vertexJsonCredentials, undefined),
|
||||
}),
|
||||
})
|
||||
} else if (this.options.vertexKeyFile) {
|
||||
return new AnthropicVertex({
|
||||
projectId,
|
||||
region,
|
||||
googleAuth: new GoogleAuth({
|
||||
scopes: ["https://www.googleapis.com/auth/cloud-platform"],
|
||||
keyFile: this.options.vertexKeyFile,
|
||||
}),
|
||||
})
|
||||
} else {
|
||||
return new AnthropicVertex({ projectId, region })
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
override async *createMessage(
|
||||
|
|
@ -62,101 +66,118 @@ export class AnthropicVertexHandler extends BaseProvider implements SingleComple
|
|||
messages: Anthropic.Messages.MessageParam[],
|
||||
metadata?: ApiHandlerCreateMessageMetadata,
|
||||
): ApiStream {
|
||||
let {
|
||||
id,
|
||||
info: { supportsPromptCache },
|
||||
temperature,
|
||||
maxTokens,
|
||||
reasoning: thinking,
|
||||
} = this.getModel()
|
||||
|
||||
/**
|
||||
* Vertex API has specific limitations for prompt caching:
|
||||
* 1. Maximum of 4 blocks can have cache_control
|
||||
* 2. Only text blocks can be cached (images and other content types cannot)
|
||||
* 3. Cache control can only be applied to user messages, not assistant messages
|
||||
*
|
||||
* Our caching strategy:
|
||||
* - Cache the system prompt (1 block)
|
||||
* - Cache the last text block of the second-to-last user message (1 block)
|
||||
* - Cache the last text block of the last user message (1 block)
|
||||
* This ensures we stay under the 4-block limit while maintaining effective caching
|
||||
* for the most relevant context.
|
||||
*/
|
||||
const params: Anthropic.Messages.MessageCreateParamsStreaming = {
|
||||
model: id,
|
||||
max_tokens: maxTokens ?? ANTHROPIC_DEFAULT_MAX_TOKENS,
|
||||
temperature,
|
||||
thinking,
|
||||
// Cache the system prompt if caching is enabled.
|
||||
system: supportsPromptCache
|
||||
? [{ text: systemPrompt, type: "text" as const, cache_control: { type: "ephemeral" } }]
|
||||
: systemPrompt,
|
||||
messages: supportsPromptCache ? addCacheBreakpoints(messages) : messages,
|
||||
stream: true,
|
||||
// Set the environment variable before making the API call
|
||||
const originalValue = process.env.GOOGLE_CLOUD_PROJECT
|
||||
if (this.options.googleCloudProject) {
|
||||
process.env.GOOGLE_CLOUD_PROJECT = this.options.googleCloudProject
|
||||
}
|
||||
|
||||
const stream = await this.client.messages.create(params)
|
||||
try {
|
||||
let {
|
||||
id,
|
||||
info: { supportsPromptCache },
|
||||
temperature,
|
||||
maxTokens,
|
||||
reasoning: thinking,
|
||||
} = this.getModel()
|
||||
|
||||
for await (const chunk of stream) {
|
||||
switch (chunk.type) {
|
||||
case "message_start": {
|
||||
const usage = chunk.message!.usage
|
||||
/**
|
||||
* Vertex API has specific limitations for prompt caching:
|
||||
* 1. Maximum of 4 blocks can have cache_control
|
||||
* 2. Only text blocks can be cached (images and other content types cannot)
|
||||
* 3. Cache control can only be applied to user messages, not assistant messages
|
||||
*
|
||||
* Our caching strategy:
|
||||
* - Cache the system prompt (1 block)
|
||||
* - Cache the last text block of the second-to-last user message (1 block)
|
||||
* - Cache the last text block of the last user message (1 block)
|
||||
* This ensures we stay under the 4-block limit while maintaining effective caching
|
||||
* for the most relevant context.
|
||||
*/
|
||||
const params: Anthropic.Messages.MessageCreateParamsStreaming = {
|
||||
model: id,
|
||||
max_tokens: maxTokens ?? ANTHROPIC_DEFAULT_MAX_TOKENS,
|
||||
temperature,
|
||||
thinking,
|
||||
// Cache the system prompt if caching is enabled.
|
||||
system: supportsPromptCache
|
||||
? [{ text: systemPrompt, type: "text" as const, cache_control: { type: "ephemeral" } }]
|
||||
: systemPrompt,
|
||||
messages: supportsPromptCache ? addCacheBreakpoints(messages) : messages,
|
||||
stream: true,
|
||||
}
|
||||
|
||||
yield {
|
||||
type: "usage",
|
||||
inputTokens: usage.input_tokens || 0,
|
||||
outputTokens: usage.output_tokens || 0,
|
||||
cacheWriteTokens: usage.cache_creation_input_tokens || undefined,
|
||||
cacheReadTokens: usage.cache_read_input_tokens || undefined,
|
||||
const stream = await this.client.messages.create(params)
|
||||
|
||||
for await (const chunk of stream) {
|
||||
switch (chunk.type) {
|
||||
case "message_start": {
|
||||
const usage = chunk.message!.usage
|
||||
|
||||
yield {
|
||||
type: "usage",
|
||||
inputTokens: usage.input_tokens || 0,
|
||||
outputTokens: usage.output_tokens || 0,
|
||||
cacheWriteTokens: usage.cache_creation_input_tokens || undefined,
|
||||
cacheReadTokens: usage.cache_read_input_tokens || undefined,
|
||||
}
|
||||
|
||||
break
|
||||
}
|
||||
case "message_delta": {
|
||||
yield {
|
||||
type: "usage",
|
||||
inputTokens: 0,
|
||||
outputTokens: chunk.usage!.output_tokens || 0,
|
||||
}
|
||||
|
||||
break
|
||||
}
|
||||
case "message_delta": {
|
||||
yield {
|
||||
type: "usage",
|
||||
inputTokens: 0,
|
||||
outputTokens: chunk.usage!.output_tokens || 0,
|
||||
break
|
||||
}
|
||||
case "content_block_start": {
|
||||
switch (chunk.content_block!.type) {
|
||||
case "text": {
|
||||
if (chunk.index! > 0) {
|
||||
yield { type: "text", text: "\n" }
|
||||
}
|
||||
|
||||
break
|
||||
}
|
||||
case "content_block_start": {
|
||||
switch (chunk.content_block!.type) {
|
||||
case "text": {
|
||||
if (chunk.index! > 0) {
|
||||
yield { type: "text", text: "\n" }
|
||||
yield { type: "text", text: chunk.content_block!.text }
|
||||
break
|
||||
}
|
||||
case "thinking": {
|
||||
if (chunk.index! > 0) {
|
||||
yield { type: "reasoning", text: "\n" }
|
||||
}
|
||||
|
||||
yield { type: "text", text: chunk.content_block!.text }
|
||||
break
|
||||
}
|
||||
case "thinking": {
|
||||
if (chunk.index! > 0) {
|
||||
yield { type: "reasoning", text: "\n" }
|
||||
yield { type: "reasoning", text: (chunk.content_block as any).thinking }
|
||||
break
|
||||
}
|
||||
|
||||
yield { type: "reasoning", text: (chunk.content_block as any).thinking }
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
break
|
||||
break
|
||||
}
|
||||
case "content_block_delta": {
|
||||
switch (chunk.delta!.type) {
|
||||
case "text_delta": {
|
||||
yield { type: "text", text: chunk.delta!.text }
|
||||
break
|
||||
}
|
||||
case "thinking_delta": {
|
||||
yield { type: "reasoning", text: (chunk.delta as any).thinking }
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
break
|
||||
}
|
||||
}
|
||||
case "content_block_delta": {
|
||||
switch (chunk.delta!.type) {
|
||||
case "text_delta": {
|
||||
yield { type: "text", text: chunk.delta!.text }
|
||||
break
|
||||
}
|
||||
case "thinking_delta": {
|
||||
yield { type: "reasoning", text: (chunk.delta as any).thinking }
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
break
|
||||
}
|
||||
} finally {
|
||||
// Restore the original environment variable value
|
||||
if (this.options.googleCloudProject) {
|
||||
if (originalValue !== undefined) {
|
||||
process.env.GOOGLE_CLOUD_PROJECT = originalValue
|
||||
} else {
|
||||
delete process.env.GOOGLE_CLOUD_PROJECT
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -176,45 +197,47 @@ export class AnthropicVertexHandler extends BaseProvider implements SingleComple
|
|||
}
|
||||
|
||||
async completePrompt(prompt: string) {
|
||||
try {
|
||||
let {
|
||||
id,
|
||||
info: { supportsPromptCache },
|
||||
temperature,
|
||||
maxTokens = ANTHROPIC_DEFAULT_MAX_TOKENS,
|
||||
reasoning: thinking,
|
||||
} = this.getModel()
|
||||
return await withGoogleCloudProject(this.options.googleCloudProject, async () => {
|
||||
try {
|
||||
let {
|
||||
id,
|
||||
info: { supportsPromptCache },
|
||||
temperature,
|
||||
maxTokens = ANTHROPIC_DEFAULT_MAX_TOKENS,
|
||||
reasoning: thinking,
|
||||
} = this.getModel()
|
||||
|
||||
const params: Anthropic.Messages.MessageCreateParamsNonStreaming = {
|
||||
model: id,
|
||||
max_tokens: maxTokens,
|
||||
temperature,
|
||||
thinking,
|
||||
messages: [
|
||||
{
|
||||
role: "user",
|
||||
content: supportsPromptCache
|
||||
? [{ type: "text" as const, text: prompt, cache_control: { type: "ephemeral" } }]
|
||||
: prompt,
|
||||
},
|
||||
],
|
||||
stream: false,
|
||||
const params: Anthropic.Messages.MessageCreateParamsNonStreaming = {
|
||||
model: id,
|
||||
max_tokens: maxTokens,
|
||||
temperature,
|
||||
thinking,
|
||||
messages: [
|
||||
{
|
||||
role: "user",
|
||||
content: supportsPromptCache
|
||||
? [{ type: "text" as const, text: prompt, cache_control: { type: "ephemeral" } }]
|
||||
: prompt,
|
||||
},
|
||||
],
|
||||
stream: false,
|
||||
}
|
||||
|
||||
const response = await this.client.messages.create(params)
|
||||
const content = response.content[0]
|
||||
|
||||
if (content.type === "text") {
|
||||
return content.text
|
||||
}
|
||||
|
||||
return ""
|
||||
} catch (error) {
|
||||
if (error instanceof Error) {
|
||||
throw new Error(`Vertex completion error: ${error.message}`)
|
||||
}
|
||||
|
||||
throw error
|
||||
}
|
||||
|
||||
const response = await this.client.messages.create(params)
|
||||
const content = response.content[0]
|
||||
|
||||
if (content.type === "text") {
|
||||
return content.text
|
||||
}
|
||||
|
||||
return ""
|
||||
} catch (error) {
|
||||
if (error instanceof Error) {
|
||||
throw new Error(`Vertex completion error: ${error.message}`)
|
||||
}
|
||||
|
||||
throw error
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ import { type ModelInfo, type GeminiModelId, geminiDefaultModelId, geminiModels
|
|||
|
||||
import type { ApiHandlerOptions } from "../../shared/api"
|
||||
import { safeJsonParse } from "../../shared/safeJsonParse"
|
||||
import { withGoogleCloudProject, withGoogleCloudProjectSync } from "../../utils/googleCloudEnv"
|
||||
|
||||
import { convertAnthropicContentToGemini, convertAnthropicMessageToGemini } from "../transform/gemini-format"
|
||||
import type { ApiStream } from "../transform/stream"
|
||||
|
|
@ -37,25 +38,28 @@ export class GeminiHandler extends BaseProvider implements SingleCompletionHandl
|
|||
const location = this.options.vertexRegion ?? "not-provided"
|
||||
const apiKey = this.options.geminiApiKey ?? "not-provided"
|
||||
|
||||
this.client = this.options.vertexJsonCredentials
|
||||
? new GoogleGenAI({
|
||||
vertexai: true,
|
||||
project,
|
||||
location,
|
||||
googleAuthOptions: {
|
||||
credentials: safeJsonParse<JWTInput>(this.options.vertexJsonCredentials, undefined),
|
||||
},
|
||||
})
|
||||
: this.options.vertexKeyFile
|
||||
// Set GOOGLE_CLOUD_PROJECT environment variable if specified in profile
|
||||
this.client = withGoogleCloudProjectSync(this.options.googleCloudProject, () => {
|
||||
return this.options.vertexJsonCredentials
|
||||
? new GoogleGenAI({
|
||||
vertexai: true,
|
||||
project,
|
||||
location,
|
||||
googleAuthOptions: { keyFile: this.options.vertexKeyFile },
|
||||
googleAuthOptions: {
|
||||
credentials: safeJsonParse<JWTInput>(this.options.vertexJsonCredentials, undefined),
|
||||
},
|
||||
})
|
||||
: isVertex
|
||||
? new GoogleGenAI({ vertexai: true, project, location })
|
||||
: new GoogleGenAI({ apiKey })
|
||||
: this.options.vertexKeyFile
|
||||
? new GoogleGenAI({
|
||||
vertexai: true,
|
||||
project,
|
||||
location,
|
||||
googleAuthOptions: { keyFile: this.options.vertexKeyFile },
|
||||
})
|
||||
: isVertex
|
||||
? new GoogleGenAI({ vertexai: true, project, location })
|
||||
: new GoogleGenAI({ apiKey })
|
||||
})
|
||||
}
|
||||
|
||||
async *createMessage(
|
||||
|
|
@ -63,68 +67,87 @@ export class GeminiHandler extends BaseProvider implements SingleCompletionHandl
|
|||
messages: Anthropic.Messages.MessageParam[],
|
||||
metadata?: ApiHandlerCreateMessageMetadata,
|
||||
): ApiStream {
|
||||
const { id: model, info, reasoning: thinkingConfig, maxTokens } = this.getModel()
|
||||
|
||||
const contents = messages.map(convertAnthropicMessageToGemini)
|
||||
|
||||
const config: GenerateContentConfig = {
|
||||
systemInstruction,
|
||||
httpOptions: this.options.googleGeminiBaseUrl ? { baseUrl: this.options.googleGeminiBaseUrl } : undefined,
|
||||
thinkingConfig,
|
||||
maxOutputTokens: this.options.modelMaxTokens ?? maxTokens ?? undefined,
|
||||
temperature: this.options.modelTemperature ?? 0,
|
||||
// Set the environment variable before making the API call
|
||||
const originalValue = process.env.GOOGLE_CLOUD_PROJECT
|
||||
if (this.options.googleCloudProject) {
|
||||
process.env.GOOGLE_CLOUD_PROJECT = this.options.googleCloudProject
|
||||
}
|
||||
|
||||
const params: GenerateContentParameters = { model, contents, config }
|
||||
try {
|
||||
const { id: model, info, reasoning: thinkingConfig, maxTokens } = this.getModel()
|
||||
|
||||
const result = await this.client.models.generateContentStream(params)
|
||||
const contents = messages.map(convertAnthropicMessageToGemini)
|
||||
|
||||
let lastUsageMetadata: GenerateContentResponseUsageMetadata | undefined
|
||||
const config: GenerateContentConfig = {
|
||||
systemInstruction,
|
||||
httpOptions: this.options.googleGeminiBaseUrl
|
||||
? { baseUrl: this.options.googleGeminiBaseUrl }
|
||||
: undefined,
|
||||
thinkingConfig,
|
||||
maxOutputTokens: this.options.modelMaxTokens ?? maxTokens ?? undefined,
|
||||
temperature: this.options.modelTemperature ?? 0,
|
||||
}
|
||||
|
||||
for await (const chunk of result) {
|
||||
// Process candidates and their parts to separate thoughts from content
|
||||
if (chunk.candidates && chunk.candidates.length > 0) {
|
||||
const candidate = chunk.candidates[0]
|
||||
if (candidate.content && candidate.content.parts) {
|
||||
for (const part of candidate.content.parts) {
|
||||
if (part.thought) {
|
||||
// This is a thinking/reasoning part
|
||||
if (part.text) {
|
||||
yield { type: "reasoning", text: part.text }
|
||||
}
|
||||
} else {
|
||||
// This is regular content
|
||||
if (part.text) {
|
||||
yield { type: "text", text: part.text }
|
||||
const params: GenerateContentParameters = { model, contents, config }
|
||||
|
||||
const result = await this.client.models.generateContentStream(params)
|
||||
|
||||
let lastUsageMetadata: GenerateContentResponseUsageMetadata | undefined
|
||||
|
||||
for await (const chunk of result) {
|
||||
// Process candidates and their parts to separate thoughts from content
|
||||
if (chunk.candidates && chunk.candidates.length > 0) {
|
||||
const candidate = chunk.candidates[0]
|
||||
if (candidate.content && candidate.content.parts) {
|
||||
for (const part of candidate.content.parts) {
|
||||
if (part.thought) {
|
||||
// This is a thinking/reasoning part
|
||||
if (part.text) {
|
||||
yield { type: "reasoning", text: part.text }
|
||||
}
|
||||
} else {
|
||||
// This is regular content
|
||||
if (part.text) {
|
||||
yield { type: "text", text: part.text }
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Fallback to the original text property if no candidates structure
|
||||
else if (chunk.text) {
|
||||
yield { type: "text", text: chunk.text }
|
||||
}
|
||||
|
||||
if (chunk.usageMetadata) {
|
||||
lastUsageMetadata = chunk.usageMetadata
|
||||
}
|
||||
}
|
||||
|
||||
// Fallback to the original text property if no candidates structure
|
||||
else if (chunk.text) {
|
||||
yield { type: "text", text: chunk.text }
|
||||
if (lastUsageMetadata) {
|
||||
const inputTokens = lastUsageMetadata.promptTokenCount ?? 0
|
||||
const outputTokens = lastUsageMetadata.candidatesTokenCount ?? 0
|
||||
const cacheReadTokens = lastUsageMetadata.cachedContentTokenCount
|
||||
const reasoningTokens = lastUsageMetadata.thoughtsTokenCount
|
||||
|
||||
yield {
|
||||
type: "usage",
|
||||
inputTokens,
|
||||
outputTokens,
|
||||
cacheReadTokens,
|
||||
reasoningTokens,
|
||||
totalCost: this.calculateCost({ info, inputTokens, outputTokens, cacheReadTokens }),
|
||||
}
|
||||
}
|
||||
|
||||
if (chunk.usageMetadata) {
|
||||
lastUsageMetadata = chunk.usageMetadata
|
||||
}
|
||||
}
|
||||
|
||||
if (lastUsageMetadata) {
|
||||
const inputTokens = lastUsageMetadata.promptTokenCount ?? 0
|
||||
const outputTokens = lastUsageMetadata.candidatesTokenCount ?? 0
|
||||
const cacheReadTokens = lastUsageMetadata.cachedContentTokenCount
|
||||
const reasoningTokens = lastUsageMetadata.thoughtsTokenCount
|
||||
|
||||
yield {
|
||||
type: "usage",
|
||||
inputTokens,
|
||||
outputTokens,
|
||||
cacheReadTokens,
|
||||
reasoningTokens,
|
||||
totalCost: this.calculateCost({ info, inputTokens, outputTokens, cacheReadTokens }),
|
||||
} finally {
|
||||
// Restore the original environment variable value
|
||||
if (this.options.googleCloudProject) {
|
||||
if (originalValue !== undefined) {
|
||||
process.env.GOOGLE_CLOUD_PROJECT = originalValue
|
||||
} else {
|
||||
delete process.env.GOOGLE_CLOUD_PROJECT
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -143,49 +166,53 @@ export class GeminiHandler extends BaseProvider implements SingleCompletionHandl
|
|||
}
|
||||
|
||||
async completePrompt(prompt: string): Promise<string> {
|
||||
try {
|
||||
const { id: model } = this.getModel()
|
||||
return await withGoogleCloudProject(this.options.googleCloudProject, async () => {
|
||||
try {
|
||||
const { id: model } = this.getModel()
|
||||
|
||||
const result = await this.client.models.generateContent({
|
||||
model,
|
||||
contents: [{ role: "user", parts: [{ text: prompt }] }],
|
||||
config: {
|
||||
httpOptions: this.options.googleGeminiBaseUrl
|
||||
? { baseUrl: this.options.googleGeminiBaseUrl }
|
||||
: undefined,
|
||||
temperature: this.options.modelTemperature ?? 0,
|
||||
},
|
||||
})
|
||||
const result = await this.client.models.generateContent({
|
||||
model,
|
||||
contents: [{ role: "user", parts: [{ text: prompt }] }],
|
||||
config: {
|
||||
httpOptions: this.options.googleGeminiBaseUrl
|
||||
? { baseUrl: this.options.googleGeminiBaseUrl }
|
||||
: undefined,
|
||||
temperature: this.options.modelTemperature ?? 0,
|
||||
},
|
||||
})
|
||||
|
||||
return result.text ?? ""
|
||||
} catch (error) {
|
||||
if (error instanceof Error) {
|
||||
throw new Error(`Gemini completion error: ${error.message}`)
|
||||
return result.text ?? ""
|
||||
} catch (error) {
|
||||
if (error instanceof Error) {
|
||||
throw new Error(`Gemini completion error: ${error.message}`)
|
||||
}
|
||||
|
||||
throw error
|
||||
}
|
||||
|
||||
throw error
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
override async countTokens(content: Array<Anthropic.Messages.ContentBlockParam>): Promise<number> {
|
||||
try {
|
||||
const { id: model } = this.getModel()
|
||||
return await withGoogleCloudProject(this.options.googleCloudProject, async () => {
|
||||
try {
|
||||
const { id: model } = this.getModel()
|
||||
|
||||
const response = await this.client.models.countTokens({
|
||||
model,
|
||||
contents: convertAnthropicContentToGemini(content),
|
||||
})
|
||||
const response = await this.client.models.countTokens({
|
||||
model,
|
||||
contents: convertAnthropicContentToGemini(content),
|
||||
})
|
||||
|
||||
if (response.totalTokens === undefined) {
|
||||
console.warn("Gemini token counting returned undefined, using fallback")
|
||||
if (response.totalTokens === undefined) {
|
||||
console.warn("Gemini token counting returned undefined, using fallback")
|
||||
return super.countTokens(content)
|
||||
}
|
||||
|
||||
return response.totalTokens
|
||||
} catch (error) {
|
||||
console.warn("Gemini token counting failed, using fallback", error)
|
||||
return super.countTokens(content)
|
||||
}
|
||||
|
||||
return response.totalTokens
|
||||
} catch (error) {
|
||||
console.warn("Gemini token counting failed, using fallback", error)
|
||||
return super.countTokens(content)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
public calculateCost({
|
||||
|
|
|
|||
137
src/utils/__tests__/googleCloudEnv.spec.ts
Normal file
137
src/utils/__tests__/googleCloudEnv.spec.ts
Normal file
|
|
@ -0,0 +1,137 @@
|
|||
// npx vitest run src/utils/__tests__/googleCloudEnv.spec.ts
|
||||
|
||||
import { describe, it, expect, beforeEach, afterEach } from "vitest"
|
||||
|
||||
import { withGoogleCloudProject, withGoogleCloudProjectSync } from "../googleCloudEnv"
|
||||
|
||||
describe("googleCloudEnv", () => {
|
||||
let originalValue: string | undefined
|
||||
|
||||
beforeEach(() => {
|
||||
// Store the original value
|
||||
originalValue = process.env.GOOGLE_CLOUD_PROJECT
|
||||
// Clean up the environment variable
|
||||
delete process.env.GOOGLE_CLOUD_PROJECT
|
||||
})
|
||||
|
||||
afterEach(() => {
|
||||
// Restore the original value
|
||||
if (originalValue !== undefined) {
|
||||
process.env.GOOGLE_CLOUD_PROJECT = originalValue
|
||||
} else {
|
||||
delete process.env.GOOGLE_CLOUD_PROJECT
|
||||
}
|
||||
})
|
||||
|
||||
describe("withGoogleCloudProject", () => {
|
||||
it("should set and restore GOOGLE_CLOUD_PROJECT environment variable", async () => {
|
||||
const testProjectId = "test-project-123"
|
||||
let capturedEnvValue: string | undefined
|
||||
|
||||
const result = await withGoogleCloudProject(testProjectId, async () => {
|
||||
capturedEnvValue = process.env.GOOGLE_CLOUD_PROJECT
|
||||
return "test-result"
|
||||
})
|
||||
|
||||
expect(capturedEnvValue).toBe(testProjectId)
|
||||
expect(process.env.GOOGLE_CLOUD_PROJECT).toBeUndefined()
|
||||
expect(result).toBe("test-result")
|
||||
})
|
||||
|
||||
it("should restore original value if it existed", async () => {
|
||||
const originalProjectId = "original-project"
|
||||
const testProjectId = "test-project-123"
|
||||
process.env.GOOGLE_CLOUD_PROJECT = originalProjectId
|
||||
|
||||
await withGoogleCloudProject(testProjectId, async () => {
|
||||
expect(process.env.GOOGLE_CLOUD_PROJECT).toBe(testProjectId)
|
||||
return "test-result"
|
||||
})
|
||||
|
||||
expect(process.env.GOOGLE_CLOUD_PROJECT).toBe(originalProjectId)
|
||||
})
|
||||
|
||||
it("should not modify environment if projectId is undefined", async () => {
|
||||
const originalProjectId = "original-project"
|
||||
process.env.GOOGLE_CLOUD_PROJECT = originalProjectId
|
||||
|
||||
await withGoogleCloudProject(undefined, async () => {
|
||||
expect(process.env.GOOGLE_CLOUD_PROJECT).toBe(originalProjectId)
|
||||
return "test-result"
|
||||
})
|
||||
|
||||
expect(process.env.GOOGLE_CLOUD_PROJECT).toBe(originalProjectId)
|
||||
})
|
||||
|
||||
it("should restore environment even if function throws", async () => {
|
||||
const testProjectId = "test-project-123"
|
||||
|
||||
try {
|
||||
await withGoogleCloudProject(testProjectId, async () => {
|
||||
expect(process.env.GOOGLE_CLOUD_PROJECT).toBe(testProjectId)
|
||||
throw new Error("Test error")
|
||||
})
|
||||
} catch (error) {
|
||||
expect(error).toBeInstanceOf(Error)
|
||||
}
|
||||
|
||||
expect(process.env.GOOGLE_CLOUD_PROJECT).toBeUndefined()
|
||||
})
|
||||
})
|
||||
|
||||
describe("withGoogleCloudProjectSync", () => {
|
||||
it("should set and restore GOOGLE_CLOUD_PROJECT environment variable synchronously", () => {
|
||||
const testProjectId = "test-project-123"
|
||||
let capturedEnvValue: string | undefined
|
||||
|
||||
const result = withGoogleCloudProjectSync(testProjectId, () => {
|
||||
capturedEnvValue = process.env.GOOGLE_CLOUD_PROJECT
|
||||
return "test-result"
|
||||
})
|
||||
|
||||
expect(capturedEnvValue).toBe(testProjectId)
|
||||
expect(process.env.GOOGLE_CLOUD_PROJECT).toBeUndefined()
|
||||
expect(result).toBe("test-result")
|
||||
})
|
||||
|
||||
it("should restore original value if it existed", () => {
|
||||
const originalProjectId = "original-project"
|
||||
const testProjectId = "test-project-123"
|
||||
process.env.GOOGLE_CLOUD_PROJECT = originalProjectId
|
||||
|
||||
withGoogleCloudProjectSync(testProjectId, () => {
|
||||
expect(process.env.GOOGLE_CLOUD_PROJECT).toBe(testProjectId)
|
||||
return "test-result"
|
||||
})
|
||||
|
||||
expect(process.env.GOOGLE_CLOUD_PROJECT).toBe(originalProjectId)
|
||||
})
|
||||
|
||||
it("should not modify environment if projectId is undefined", () => {
|
||||
const originalProjectId = "original-project"
|
||||
process.env.GOOGLE_CLOUD_PROJECT = originalProjectId
|
||||
|
||||
withGoogleCloudProjectSync(undefined, () => {
|
||||
expect(process.env.GOOGLE_CLOUD_PROJECT).toBe(originalProjectId)
|
||||
return "test-result"
|
||||
})
|
||||
|
||||
expect(process.env.GOOGLE_CLOUD_PROJECT).toBe(originalProjectId)
|
||||
})
|
||||
|
||||
it("should restore environment even if function throws", () => {
|
||||
const testProjectId = "test-project-123"
|
||||
|
||||
try {
|
||||
withGoogleCloudProjectSync(testProjectId, () => {
|
||||
expect(process.env.GOOGLE_CLOUD_PROJECT).toBe(testProjectId)
|
||||
throw new Error("Test error")
|
||||
})
|
||||
} catch (error) {
|
||||
expect(error).toBeInstanceOf(Error)
|
||||
}
|
||||
|
||||
expect(process.env.GOOGLE_CLOUD_PROJECT).toBeUndefined()
|
||||
})
|
||||
})
|
||||
})
|
||||
63
src/utils/googleCloudEnv.ts
Normal file
63
src/utils/googleCloudEnv.ts
Normal file
|
|
@ -0,0 +1,63 @@
|
|||
/**
|
||||
* Utility functions for managing Google Cloud environment variables
|
||||
* for provider-specific configurations
|
||||
*/
|
||||
|
||||
/**
|
||||
* Temporarily sets the GOOGLE_CLOUD_PROJECT environment variable
|
||||
* and executes a function, then restores the original value
|
||||
*/
|
||||
export async function withGoogleCloudProject<T>(projectId: string | undefined, fn: () => Promise<T>): Promise<T> {
|
||||
if (!projectId) {
|
||||
// If no project ID is specified, just execute the function
|
||||
return await fn()
|
||||
}
|
||||
|
||||
const originalValue = process.env.GOOGLE_CLOUD_PROJECT
|
||||
|
||||
try {
|
||||
// Set the environment variable
|
||||
process.env.GOOGLE_CLOUD_PROJECT = projectId
|
||||
|
||||
// Execute the function
|
||||
const result = await fn()
|
||||
|
||||
return result
|
||||
} finally {
|
||||
// Restore the original value
|
||||
if (originalValue !== undefined) {
|
||||
process.env.GOOGLE_CLOUD_PROJECT = originalValue
|
||||
} else {
|
||||
delete process.env.GOOGLE_CLOUD_PROJECT
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Synchronous version for non-async operations
|
||||
*/
|
||||
export function withGoogleCloudProjectSync<T>(projectId: string | undefined, fn: () => T): T {
|
||||
if (!projectId) {
|
||||
// If no project ID is specified, just execute the function
|
||||
return fn()
|
||||
}
|
||||
|
||||
const originalValue = process.env.GOOGLE_CLOUD_PROJECT
|
||||
|
||||
try {
|
||||
// Set the environment variable
|
||||
process.env.GOOGLE_CLOUD_PROJECT = projectId
|
||||
|
||||
// Execute the function
|
||||
const result = fn()
|
||||
|
||||
return result
|
||||
} finally {
|
||||
// Restore the original value
|
||||
if (originalValue !== undefined) {
|
||||
process.env.GOOGLE_CLOUD_PROJECT = originalValue
|
||||
} else {
|
||||
delete process.env.GOOGLE_CLOUD_PROJECT
|
||||
}
|
||||
}
|
||||
}
|
||||
Loading…
Add table
Reference in a new issue