@@ -0,0 +1,501 @@
|
||||
import {
|
||||
buildChatToolSystemPrompt,
|
||||
executeToolCallAndBuildEvent,
|
||||
getEnabledChatTools,
|
||||
getUnstreamedText,
|
||||
looksLikeDanglingToolIntent,
|
||||
MAX_DANGLING_TOOL_INTENT_RETRIES,
|
||||
MAX_TOOL_ROUNDS,
|
||||
prepareToolCallExecution,
|
||||
type NormalizedToolCall,
|
||||
type ToolAwareCompletionParams,
|
||||
type ToolAwareCompletionResult,
|
||||
type ToolAwareStreamingEvent,
|
||||
type ToolAwareUsage,
|
||||
type ToolExecutionEvent,
|
||||
} from "../chat-tools.js";
|
||||
import {
|
||||
buildImageSummaryText,
|
||||
buildTextAttachmentPrompt,
|
||||
buildTopLevelSystemPrompt,
|
||||
getImageAttachments,
|
||||
getTextAttachments,
|
||||
parseImageDataUrl,
|
||||
} from "../message-content.js";
|
||||
import type { ChatMessage } from "../types.js";
|
||||
|
||||
type GeminiClient = {
|
||||
apiKey: string;
|
||||
baseURL: string;
|
||||
};
|
||||
|
||||
const INTERNAL_CORRECTION =
|
||||
"Internal correction: the previous assistant message claimed it would run a tool, but no tool call was made. If the task needs an available tool, call it now. Otherwise provide the final answer directly without saying you will run a tool.";
|
||||
|
||||
function normalizeModelResourceName(model: string) {
|
||||
const trimmed = model.trim().replace(/^\/+/, "");
|
||||
return trimmed.startsWith("models/") || trimmed.startsWith("tunedModels/") ? trimmed : `models/${trimmed}`;
|
||||
}
|
||||
|
||||
function geminiUrl(client: GeminiClient, model: string, method: "generateContent" | "streamGenerateContent", extraParams: Record<string, string> = {}) {
|
||||
const url = new URL(`${client.baseURL.replace(/\/+$/, "")}/${normalizeModelResourceName(model)}:${method}`);
|
||||
url.searchParams.set("key", client.apiKey);
|
||||
for (const [key, value] of Object.entries(extraParams)) {
|
||||
url.searchParams.set(key, value);
|
||||
}
|
||||
return url;
|
||||
}
|
||||
|
||||
function generationConfig(params: Pick<ToolAwareCompletionParams, "temperature" | "maxTokens">) {
|
||||
const config: Record<string, unknown> = {};
|
||||
if (params.temperature !== undefined) config.temperature = params.temperature;
|
||||
if (params.maxTokens !== undefined) config.maxOutputTokens = params.maxTokens;
|
||||
return Object.keys(config).length ? config : undefined;
|
||||
}
|
||||
|
||||
function toGeminiJsonSchema(schema: unknown): Record<string, unknown> | undefined {
|
||||
if (!schema || typeof schema !== "object" || Array.isArray(schema)) return undefined;
|
||||
const input = schema as Record<string, unknown>;
|
||||
const output: Record<string, unknown> = {};
|
||||
|
||||
if (typeof input.type === "string") output.type = input.type;
|
||||
if (typeof input.description === "string") output.description = input.description;
|
||||
if (typeof input.format === "string") output.format = input.format;
|
||||
if (typeof input.nullable === "boolean") output.nullable = input.nullable;
|
||||
if (Array.isArray(input.enum)) output.enum = input.enum.filter((value) => typeof value === "string");
|
||||
if (Array.isArray(input.required)) output.required = input.required.filter((value) => typeof value === "string");
|
||||
|
||||
const items = toGeminiJsonSchema(input.items);
|
||||
if (items) output.items = items;
|
||||
|
||||
if (input.properties && typeof input.properties === "object" && !Array.isArray(input.properties)) {
|
||||
const properties: Record<string, unknown> = {};
|
||||
for (const [key, value] of Object.entries(input.properties)) {
|
||||
const propertySchema = toGeminiJsonSchema(value);
|
||||
if (propertySchema) properties[key] = propertySchema;
|
||||
}
|
||||
if (Object.keys(properties).length) output.properties = properties;
|
||||
}
|
||||
|
||||
return Object.keys(output).length ? output : undefined;
|
||||
}
|
||||
|
||||
function toGeminiTools(tools: any[]) {
|
||||
const functionDeclarations = tools
|
||||
.map((tool) => {
|
||||
if (tool?.type !== "function") return null;
|
||||
const declaration: Record<string, unknown> = {
|
||||
name: tool.function.name,
|
||||
description: tool.function.description,
|
||||
};
|
||||
const parameters = toGeminiJsonSchema(tool.function.parameters);
|
||||
if (parameters) declaration.parameters = parameters;
|
||||
return declaration;
|
||||
})
|
||||
.filter(Boolean);
|
||||
|
||||
return functionDeclarations.length ? [{ functionDeclarations }] : undefined;
|
||||
}
|
||||
|
||||
function toContentParts(message: ChatMessage) {
|
||||
const imageAttachments = getImageAttachments(message);
|
||||
const textAttachments = getTextAttachments(message);
|
||||
const parts: Array<Record<string, unknown>> = [];
|
||||
|
||||
for (const attachment of imageAttachments) {
|
||||
const source = parseImageDataUrl(attachment);
|
||||
parts.push({
|
||||
inlineData: {
|
||||
mimeType: source.mediaType,
|
||||
data: source.data,
|
||||
},
|
||||
});
|
||||
}
|
||||
|
||||
const imageSummary = buildImageSummaryText(imageAttachments);
|
||||
if (imageSummary) {
|
||||
parts.push({ text: imageSummary });
|
||||
}
|
||||
|
||||
for (const attachment of textAttachments) {
|
||||
parts.push({ text: buildTextAttachmentPrompt(attachment) });
|
||||
}
|
||||
|
||||
if (message.content.trim()) {
|
||||
parts.push({ text: message.content });
|
||||
}
|
||||
|
||||
return parts.length ? parts : [{ text: "" }];
|
||||
}
|
||||
|
||||
function buildConversationContent(message: ChatMessage) {
|
||||
if (message.role === "system") {
|
||||
throw new Error("System messages must be handled separately for Gemini.");
|
||||
}
|
||||
|
||||
if (message.role === "tool") {
|
||||
const name = message.name?.trim() || "tool";
|
||||
return {
|
||||
role: "user",
|
||||
parts: [{ text: `Tool output (${name}):\n${message.content}` }],
|
||||
};
|
||||
}
|
||||
|
||||
return {
|
||||
role: message.role === "assistant" ? "model" : "user",
|
||||
parts: toContentParts(message),
|
||||
};
|
||||
}
|
||||
|
||||
function buildBaseContents(messages: ChatMessage[]) {
|
||||
return messages.filter((message) => message.role !== "system").map((message) => buildConversationContent(message));
|
||||
}
|
||||
|
||||
function buildSystemInstruction(params: ToolAwareCompletionParams, toolSystemPrompt?: string) {
|
||||
const text = buildTopLevelSystemPrompt(params.messages, params.userLocation, toolSystemPrompt);
|
||||
return text ? { parts: [{ text }] } : undefined;
|
||||
}
|
||||
|
||||
function mergeUsage(acc: Required<ToolAwareUsage>, usage: any) {
|
||||
const normalized = normalizeUsage(usage);
|
||||
if (!normalized) return false;
|
||||
acc.inputTokens += normalized.inputTokens;
|
||||
acc.outputTokens += normalized.outputTokens;
|
||||
acc.totalTokens += normalized.totalTokens;
|
||||
return true;
|
||||
}
|
||||
|
||||
function normalizeUsage(usage: any) {
|
||||
if (!usage) return null;
|
||||
const inputTokens = usage.promptTokenCount ?? 0;
|
||||
const outputTokens = usage.candidatesTokenCount ?? 0;
|
||||
const totalTokens = usage.totalTokenCount ?? inputTokens + outputTokens;
|
||||
return { inputTokens, outputTokens, totalTokens };
|
||||
}
|
||||
|
||||
function getCandidate(response: any) {
|
||||
return Array.isArray(response?.candidates) ? response.candidates[0] : null;
|
||||
}
|
||||
|
||||
function getParts(response: any) {
|
||||
const parts = getCandidate(response)?.content?.parts;
|
||||
return Array.isArray(parts) ? parts : [];
|
||||
}
|
||||
|
||||
function extractText(response: any) {
|
||||
return getParts(response)
|
||||
.map((part: any) => (typeof part?.text === "string" ? part.text : ""))
|
||||
.join("");
|
||||
}
|
||||
|
||||
function stringifyToolArgs(args: unknown) {
|
||||
try {
|
||||
return JSON.stringify(args ?? {});
|
||||
} catch {
|
||||
return "{}";
|
||||
}
|
||||
}
|
||||
|
||||
function normalizeToolCallsFromParts(parts: any[], round: number): NormalizedToolCall[] {
|
||||
return parts
|
||||
.filter((part) => part?.functionCall)
|
||||
.map((part, index) => ({
|
||||
id: part.functionCall.id ?? `tool_call_${round}_${index}`,
|
||||
name: part.functionCall.name ?? "unknown_tool",
|
||||
arguments: stringifyToolArgs(part.functionCall.args),
|
||||
}));
|
||||
}
|
||||
|
||||
function buildFunctionResponsePart(call: NormalizedToolCall, toolResult: unknown) {
|
||||
return {
|
||||
functionResponse: {
|
||||
id: call.id,
|
||||
name: call.name,
|
||||
response: toolResult,
|
||||
},
|
||||
};
|
||||
}
|
||||
|
||||
function appendCorrection(conversation: any[], text: string) {
|
||||
conversation.push({ role: "model", parts: [{ text }] });
|
||||
conversation.push({ role: "user", parts: [{ text: INTERNAL_CORRECTION }] });
|
||||
}
|
||||
|
||||
async function parseGeminiResponse(response: Response) {
|
||||
const bodyText = await response.text();
|
||||
let body: any = null;
|
||||
try {
|
||||
body = bodyText ? JSON.parse(bodyText) : null;
|
||||
} catch {
|
||||
body = { raw: bodyText };
|
||||
}
|
||||
|
||||
if (!response.ok) {
|
||||
throw new Error(body?.error?.message ?? `Gemini API request failed with status ${response.status}.`);
|
||||
}
|
||||
|
||||
return body;
|
||||
}
|
||||
|
||||
async function generateContent(params: ToolAwareCompletionParams, body: Record<string, unknown>) {
|
||||
const response = await fetch(geminiUrl(params.client, params.model, "generateContent"), {
|
||||
method: "POST",
|
||||
headers: { "Content-Type": "application/json" },
|
||||
body: JSON.stringify(body),
|
||||
});
|
||||
return parseGeminiResponse(response);
|
||||
}
|
||||
|
||||
function getFailureMessage(response: any, text: string, toolCallCount: number) {
|
||||
const promptBlockReason = response?.promptFeedback?.blockReason;
|
||||
if (promptBlockReason) return `Gemini prompt blocked: ${promptBlockReason}.`;
|
||||
|
||||
const candidate = getCandidate(response);
|
||||
const finishReason = candidate?.finishReason;
|
||||
if (!finishReason || finishReason === "STOP" || finishReason === "MAX_TOKENS") return null;
|
||||
if (text || toolCallCount > 0) return null;
|
||||
return candidate?.finishMessage ?? `Gemini response stopped: ${finishReason}.`;
|
||||
}
|
||||
|
||||
function buildRequest(params: ToolAwareCompletionParams, conversation: any[], enabledTools: any[] = []) {
|
||||
const tools = toGeminiTools(enabledTools);
|
||||
return {
|
||||
contents: conversation,
|
||||
systemInstruction: buildSystemInstruction(params, enabledTools.length ? buildChatToolSystemPrompt(params) : undefined),
|
||||
generationConfig: generationConfig(params),
|
||||
tools,
|
||||
toolConfig: tools ? { functionCallingConfig: { mode: "AUTO" } } : undefined,
|
||||
};
|
||||
}
|
||||
|
||||
export async function completeWithGeminiApi(params: ToolAwareCompletionParams): Promise<ToolAwareCompletionResult> {
|
||||
const enabledTools = getEnabledChatTools(params);
|
||||
const conversation = buildBaseContents(params.messages);
|
||||
const rawResponses: unknown[] = [];
|
||||
const toolEvents: ToolExecutionEvent[] = [];
|
||||
const usageAcc: Required<ToolAwareUsage> = { inputTokens: 0, outputTokens: 0, totalTokens: 0 };
|
||||
let sawUsage = false;
|
||||
let totalToolCalls = 0;
|
||||
let danglingToolIntentRetries = 0;
|
||||
|
||||
for (let round = 0; round < MAX_TOOL_ROUNDS; round += 1) {
|
||||
const response = await generateContent(params, buildRequest(params, conversation, enabledTools));
|
||||
rawResponses.push(response);
|
||||
sawUsage = mergeUsage(usageAcc, response?.usageMetadata) || sawUsage;
|
||||
|
||||
const parts = getParts(response);
|
||||
const text = extractText(response);
|
||||
const normalizedToolCalls = normalizeToolCallsFromParts(parts, round);
|
||||
const failureMessage = getFailureMessage(response, text, normalizedToolCalls.length);
|
||||
if (failureMessage) throw new Error(failureMessage);
|
||||
|
||||
if (!normalizedToolCalls.length) {
|
||||
if (danglingToolIntentRetries < MAX_DANGLING_TOOL_INTENT_RETRIES && looksLikeDanglingToolIntent(text)) {
|
||||
danglingToolIntentRetries += 1;
|
||||
appendCorrection(conversation, text);
|
||||
continue;
|
||||
}
|
||||
return {
|
||||
text,
|
||||
usage: sawUsage ? usageAcc : undefined,
|
||||
raw: { responses: rawResponses, toolCallsUsed: totalToolCalls, api: "gemini.generateContent" },
|
||||
toolEvents,
|
||||
};
|
||||
}
|
||||
|
||||
totalToolCalls += normalizedToolCalls.length;
|
||||
conversation.push({ role: "model", parts });
|
||||
|
||||
const toolResultParts: any[] = [];
|
||||
for (const call of normalizedToolCalls) {
|
||||
const { execution } = prepareToolCallExecution(call);
|
||||
const { event, toolResult } = await executeToolCallAndBuildEvent(call, execution, params);
|
||||
toolEvents.push(event);
|
||||
toolResultParts.push(buildFunctionResponsePart(call, toolResult));
|
||||
}
|
||||
|
||||
conversation.push({ role: "user", parts: toolResultParts });
|
||||
}
|
||||
|
||||
return {
|
||||
text: "I reached the tool-call limit while gathering information. Please narrow the request and try again.",
|
||||
usage: sawUsage ? usageAcc : undefined,
|
||||
raw: { responses: rawResponses, toolCallsUsed: totalToolCalls, toolCallLimitReached: true, api: "gemini.generateContent" },
|
||||
toolEvents,
|
||||
};
|
||||
}
|
||||
|
||||
function findSseBoundary(buffer: string) {
|
||||
const crlf = buffer.indexOf("\r\n\r\n");
|
||||
const lf = buffer.indexOf("\n\n");
|
||||
if (crlf === -1) return lf === -1 ? null : { index: lf, length: 2 };
|
||||
if (lf === -1) return { index: crlf, length: 4 };
|
||||
return crlf < lf ? { index: crlf, length: 4 } : { index: lf, length: 2 };
|
||||
}
|
||||
|
||||
function parseSseEvent(rawEvent: string) {
|
||||
const data = rawEvent
|
||||
.split(/\r?\n/)
|
||||
.filter((line) => line.startsWith("data:"))
|
||||
.map((line) => line.slice("data:".length).trimStart())
|
||||
.join("\n")
|
||||
.trim();
|
||||
if (!data || data === "[DONE]") return null;
|
||||
return JSON.parse(data);
|
||||
}
|
||||
|
||||
async function* streamGeminiResponses(params: ToolAwareCompletionParams, body: Record<string, unknown>) {
|
||||
const response = await fetch(geminiUrl(params.client, params.model, "streamGenerateContent", { alt: "sse" }), {
|
||||
method: "POST",
|
||||
headers: { "Content-Type": "application/json" },
|
||||
body: JSON.stringify(body),
|
||||
});
|
||||
|
||||
if (!response.ok) {
|
||||
await parseGeminiResponse(response);
|
||||
return;
|
||||
}
|
||||
|
||||
if (!response.body) {
|
||||
throw new Error("Gemini stream response did not include a body.");
|
||||
}
|
||||
|
||||
const reader = response.body.getReader();
|
||||
const decoder = new TextDecoder();
|
||||
let buffer = "";
|
||||
|
||||
while (true) {
|
||||
const { value, done } = await reader.read();
|
||||
if (done) break;
|
||||
buffer += decoder.decode(value, { stream: true });
|
||||
let boundary = findSseBoundary(buffer);
|
||||
while (boundary) {
|
||||
const rawEvent = buffer.slice(0, boundary.index);
|
||||
buffer = buffer.slice(boundary.index + boundary.length);
|
||||
const event = parseSseEvent(rawEvent);
|
||||
if (event) yield event;
|
||||
boundary = findSseBoundary(buffer);
|
||||
}
|
||||
}
|
||||
|
||||
buffer += decoder.decode();
|
||||
const tail = buffer.trim();
|
||||
if (tail) {
|
||||
const event = parseSseEvent(tail);
|
||||
if (event) yield event;
|
||||
}
|
||||
}
|
||||
|
||||
export async function* streamWithGeminiApi(params: ToolAwareCompletionParams): AsyncGenerator<ToolAwareStreamingEvent> {
|
||||
const enabledTools = getEnabledChatTools(params);
|
||||
const conversation = buildBaseContents(params.messages);
|
||||
const rawResponses: unknown[] = [];
|
||||
const toolEvents: ToolExecutionEvent[] = [];
|
||||
const usageAcc: Required<ToolAwareUsage> = { inputTokens: 0, outputTokens: 0, totalTokens: 0 };
|
||||
let sawUsage = false;
|
||||
let totalToolCalls = 0;
|
||||
let danglingToolIntentRetries = 0;
|
||||
|
||||
if (!enabledTools.length) {
|
||||
let text = "";
|
||||
let latestUsage: any = null;
|
||||
for await (const response of streamGeminiResponses(params, buildRequest(params, conversation))) {
|
||||
rawResponses.push(response);
|
||||
if (response?.usageMetadata) latestUsage = response.usageMetadata;
|
||||
const failureMessage = getFailureMessage(response, extractText(response), 0);
|
||||
if (failureMessage) throw new Error(failureMessage);
|
||||
const delta = extractText(response);
|
||||
if (delta) {
|
||||
text += delta;
|
||||
yield { type: "delta", text: delta };
|
||||
}
|
||||
}
|
||||
|
||||
sawUsage = mergeUsage(usageAcc, latestUsage) || sawUsage;
|
||||
|
||||
yield {
|
||||
type: "done",
|
||||
result: {
|
||||
text,
|
||||
usage: sawUsage ? usageAcc : undefined,
|
||||
raw: { streamed: true, responses: rawResponses, toolCallsUsed: 0, api: "gemini.streamGenerateContent" },
|
||||
toolEvents: [],
|
||||
},
|
||||
};
|
||||
return;
|
||||
}
|
||||
|
||||
for (let round = 0; round < MAX_TOOL_ROUNDS; round += 1) {
|
||||
const roundParts: any[] = [];
|
||||
let roundText = "";
|
||||
let latestRoundResponse: any = null;
|
||||
let latestRoundUsage: any = null;
|
||||
|
||||
for await (const response of streamGeminiResponses(params, buildRequest(params, conversation, enabledTools))) {
|
||||
rawResponses.push(response);
|
||||
latestRoundResponse = response;
|
||||
if (response?.usageMetadata) latestRoundUsage = response.usageMetadata;
|
||||
roundParts.push(...getParts(response));
|
||||
roundText += extractText(response);
|
||||
}
|
||||
|
||||
sawUsage = mergeUsage(usageAcc, latestRoundUsage) || sawUsage;
|
||||
|
||||
const normalizedToolCalls = normalizeToolCallsFromParts(roundParts, round);
|
||||
const failureMessage = getFailureMessage(latestRoundResponse ?? { candidates: [{ content: { parts: roundParts } }] }, roundText, normalizedToolCalls.length);
|
||||
if (failureMessage) throw new Error(failureMessage);
|
||||
|
||||
if (!normalizedToolCalls.length) {
|
||||
if (danglingToolIntentRetries < MAX_DANGLING_TOOL_INTENT_RETRIES && looksLikeDanglingToolIntent(roundText)) {
|
||||
danglingToolIntentRetries += 1;
|
||||
appendCorrection(conversation, roundText);
|
||||
continue;
|
||||
}
|
||||
const unstreamedText = getUnstreamedText(roundText, "");
|
||||
if (unstreamedText) {
|
||||
yield { type: "delta", text: unstreamedText };
|
||||
}
|
||||
yield {
|
||||
type: "done",
|
||||
result: {
|
||||
text: roundText,
|
||||
usage: sawUsage ? usageAcc : undefined,
|
||||
raw: { streamed: true, responses: rawResponses, toolCallsUsed: totalToolCalls, api: "gemini.streamGenerateContent" },
|
||||
toolEvents,
|
||||
},
|
||||
};
|
||||
return;
|
||||
}
|
||||
|
||||
totalToolCalls += normalizedToolCalls.length;
|
||||
conversation.push({ role: "model", parts: roundParts });
|
||||
|
||||
const toolResultParts: any[] = [];
|
||||
for (const call of normalizedToolCalls) {
|
||||
const { event: initiatedEvent, execution } = prepareToolCallExecution(call);
|
||||
yield { type: "tool_call", event: initiatedEvent };
|
||||
const { event, toolResult } = await executeToolCallAndBuildEvent(call, execution, params);
|
||||
toolEvents.push(event);
|
||||
yield { type: "tool_call", event };
|
||||
toolResultParts.push(buildFunctionResponsePart(call, toolResult));
|
||||
}
|
||||
|
||||
conversation.push({ role: "user", parts: toolResultParts });
|
||||
}
|
||||
|
||||
yield {
|
||||
type: "done",
|
||||
result: {
|
||||
text: "I reached the tool-call limit while gathering information. Please narrow the request and try again.",
|
||||
usage: sawUsage ? usageAcc : undefined,
|
||||
raw: {
|
||||
streamed: true,
|
||||
responses: rawResponses,
|
||||
toolCallsUsed: totalToolCalls,
|
||||
toolCallLimitReached: true,
|
||||
api: "gemini.streamGenerateContent",
|
||||
},
|
||||
toolEvents,
|
||||
},
|
||||
};
|
||||
}
|
||||
Reference in New Issue
Block a user