diff --git a/.changeset/late-geese-itch.md b/.changeset/late-geese-itch.md new file mode 100644 index 0000000000..0c1aa30dc4 --- /dev/null +++ b/.changeset/late-geese-itch.md @@ -0,0 +1,5 @@ +--- +"claude-dev": patch +--- + +Vertex AI Gemini Flash Support diff --git a/package.json b/package.json index e5d1fdc95e..9a03f1460d 100644 --- a/package.json +++ b/package.json @@ -265,6 +265,7 @@ "@anthropic-ai/bedrock-sdk": "^0.12.4", "@anthropic-ai/sdk": "^0.37.0", "@anthropic-ai/vertex-sdk": "^0.6.4", + "@google-cloud/vertexai": "^1.9.3", "@google/generative-ai": "^0.18.0", "@mistralai/mistralai": "^1.5.0", "@modelcontextprotocol/sdk": "^1.0.1", diff --git a/src/api/providers/vertex.ts b/src/api/providers/vertex.ts index c88ceb5028..6ef9824875 100644 --- a/src/api/providers/vertex.ts +++ b/src/api/providers/vertex.ts @@ -4,19 +4,25 @@ import { withRetry } from "../retry" import { ApiHandler } from "../" import { ApiHandlerOptions, ModelInfo, vertexDefaultModelId, VertexModelId, vertexModels } from "../../shared/api" import { ApiStream } from "../transform/stream" +import { VertexAI } from "@google-cloud/vertexai" // https://docs.anthropic.com/en/api/claude-on-vertex-ai export class VertexHandler implements ApiHandler { private options: ApiHandlerOptions - private client: AnthropicVertex + private clientAnthropic: AnthropicVertex + private clientVertex: VertexAI constructor(options: ApiHandlerOptions) { this.options = options - this.client = new AnthropicVertex({ + this.clientAnthropic = new AnthropicVertex({ projectId: this.options.vertexProjectId, // https://cloud.google.com/vertex-ai/generative-ai/docs/partner-models/use-claude#regions region: this.options.vertexRegion, }) + this.clientVertex = new VertexAI({ + project: this.options.vertexProjectId, + location: this.options.vertexRegion, + }) } @withRetry() @@ -24,40 +30,66 @@ export class VertexHandler implements ApiHandler { const model = this.getModel() const modelId = model.id - let budget_tokens = this.options.thinkingBudgetTokens || 0 - const reasoningOn = budget_tokens !== 0 ? true : false + if (modelId.includes("claude")) { + let budget_tokens = this.options.thinkingBudgetTokens || 0 + const reasoningOn = budget_tokens !== 0 ? true : false - let stream - switch (modelId) { - case "claude-3-7-sonnet@20250219": - case "claude-3-5-sonnet-v2@20241022": - case "claude-3-5-sonnet@20240620": - case "claude-3-5-haiku@20241022": - case "claude-3-opus@20240229": - case "claude-3-haiku@20240307": { - // Find indices of user messages for cache control - const userMsgIndices = messages.reduce( - (acc, msg, index) => (msg.role === "user" ? [...acc, index] : acc), - [] as number[], - ) - const lastUserMsgIndex = userMsgIndices[userMsgIndices.length - 1] ?? -1 - const secondLastMsgUserIndex = userMsgIndices[userMsgIndices.length - 2] ?? -1 + let stream + switch (modelId) { + case "claude-3-7-sonnet@20250219": + case "claude-3-5-sonnet-v2@20241022": + case "claude-3-5-sonnet@20240620": + case "claude-3-5-haiku@20241022": + case "claude-3-opus@20240229": + case "claude-3-haiku@20240307": { + // Find indices of user messages for cache control + const userMsgIndices = messages.reduce( + (acc, msg, index) => (msg.role === "user" ? [...acc, index] : acc), + [] as number[], + ) + const lastUserMsgIndex = userMsgIndices[userMsgIndices.length - 1] ?? -1 + const secondLastMsgUserIndex = userMsgIndices[userMsgIndices.length - 2] ?? -1 - stream = await this.client.beta.messages.create( - { - model: modelId, - max_tokens: model.info.maxTokens || 8192, - thinking: reasoningOn ? { type: "enabled", budget_tokens: budget_tokens } : undefined, - temperature: reasoningOn ? undefined : 0, - system: [ - { - text: systemPrompt, - type: "text", - cache_control: { type: "ephemeral" }, - }, - ], - messages: messages.map((message, index) => { - if (index === lastUserMsgIndex || index === secondLastMsgUserIndex) { + stream = await this.clientAnthropic.beta.messages.create( + { + model: modelId, + max_tokens: model.info.maxTokens || 8192, + thinking: reasoningOn ? { type: "enabled", budget_tokens: budget_tokens } : undefined, + temperature: reasoningOn ? undefined : 0, + system: [ + { + text: systemPrompt, + type: "text", + cache_control: { type: "ephemeral" }, + }, + ], + messages: messages.map((message, index) => { + if (index === lastUserMsgIndex || index === secondLastMsgUserIndex) { + return { + ...message, + content: + typeof message.content === "string" + ? [ + { + type: "text", + text: message.content, + cache_control: { + type: "ephemeral", + }, + }, + ] + : message.content.map((content, contentIndex) => + contentIndex === message.content.length - 1 + ? { + ...content, + cache_control: { + type: "ephemeral", + }, + } + : content, + ), + } + } return { ...message, content: @@ -66,121 +98,152 @@ export class VertexHandler implements ApiHandler { { type: "text", text: message.content, - cache_control: { - type: "ephemeral", - }, }, ] - : message.content.map((content, contentIndex) => - contentIndex === message.content.length - 1 - ? { - ...content, - cache_control: { - type: "ephemeral", - }, - } - : content, - ), + : message.content, } - } - return { - ...message, - content: - typeof message.content === "string" - ? [ - { - type: "text", - text: message.content, - }, - ] - : message.content, - } - }), - stream: true, - }, - { - headers: {}, - }, - ) - break - } - default: { - stream = await this.client.beta.messages.create({ - model: modelId, - max_tokens: model.info.maxTokens || 8192, - temperature: 0, - system: [ - { - text: systemPrompt, - type: "text", + }), + stream: true, }, - ], - messages: messages.map((message) => ({ - ...message, - content: - typeof message.content === "string" - ? [ - { - type: "text", - text: message.content, - }, - ] - : message.content, - })), - stream: true, - }) - break + { + headers: {}, + }, + ) + break + } + default: { + stream = await this.clientAnthropic.beta.messages.create({ + model: modelId, + max_tokens: model.info.maxTokens || 8192, + temperature: 0, + system: [ + { + text: systemPrompt, + type: "text", + }, + ], + messages: messages.map((message) => ({ + ...message, + content: + typeof message.content === "string" + ? [ + { + type: "text", + text: message.content, + }, + ] + : message.content, + })), + stream: true, + }) + break + } } - } - 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_stop": - break - case "content_block_start": - switch (chunk.content_block.type) { - case "text": - if (chunk.index > 0) { + 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_stop": + 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: "\n", + text: chunk.content_block.text, } + break + } + break + case "content_block_delta": + switch (chunk.delta.type) { + case "text_delta": + yield { + type: "text", + text: chunk.delta.text, + } + break + } + break + case "content_block_stop": + break + } + } + } else { + // gemini + const generativeModel = this.clientVertex.getGenerativeModel({ + model: this.getModel().id, + systemInstruction: { + role: "system", + parts: [{ text: systemPrompt }], + }, + }) + const request = { + contents: [ + { + role: "user", + parts: messages.map((m) => { + if (typeof m.content === "string") { + return { text: m.content } + } else if (Array.isArray(m.content)) { + return { + text: m.content + .map((block) => { + if (typeof block === "string") { + return block + } else if (block.type === "text") { + return block.text + } else { + console.log("Unsupported block type", block) + return "" + } + }) + .join(" "), + } + } else { + return { text: "" } } + }), + }, + ], + } + const streamingResult = await generativeModel.generateContentStream(request) + for await (const chunk of streamingResult.stream) { + // If usage data is available, yield it similarly: + // yield { type: "usage", inputTokens: 0, outputTokens: 0 } + // Otherwise, just yield text: + const candidates = chunk.candidates || [] + for (const candidate of candidates) { + for (const part of candidate.content?.parts || []) { + if (part.text) { yield { type: "text", - text: chunk.content_block.text, + text: part.text, } - break + } } - break - case "content_block_delta": - switch (chunk.delta.type) { - case "text_delta": - yield { - type: "text", - text: chunk.delta.text, - } - break - } - break - case "content_block_stop": - break + } } } } diff --git a/src/shared/api.ts b/src/shared/api.ts index 5676584e7e..58ce768f41 100644 --- a/src/shared/api.ts +++ b/src/shared/api.ts @@ -296,6 +296,14 @@ export const vertexModels = { cacheWritesPrice: 0.3, cacheReadsPrice: 0.03, }, + "gemini-2.0-flash-001": { + maxTokens: 8192, + contextWindow: 1_048_576, + supportsImages: true, + supportsPromptCache: false, + inputPrice: 0.1, + outputPrice: 0.4, + }, } as const satisfies Record export const openAiModelInfoSaneDefaults: ModelInfo = {