mirror of
https://github.com/toeverything/AFFiNE.git
synced 2026-08-09 05:05:52 +08:00
feat(server): migrate copilot to native (#14620)
#### PR Dependency Tree * **PR #14620** 👈 This tree was auto-generated by [Charcoal](https://github.com/danerwilliams/charcoal) <!-- This is an auto-generated comment: release notes by coderabbit.ai --> ## Summary by CodeRabbit * **New Features** * Native LLM workflows: structured outputs, embeddings, and reranking plus richer multimodal attachments (images, audio, files) and improved remote-attachment inlining. * **Refactor** * Tooling API unified behind a local tool-definition helper; provider/adapters reorganized to route through native dispatch paths. * **Chores** * Dependency updates, removed legacy Google SDK integrations, and increased front memory allocation. * **Tests** * Expanded end-to-end and streaming tests exercising native provider flows, attachments, and rerank/structured scenarios. <!-- end of auto-generated comment: release notes by coderabbit.ai -->
This commit is contained in:
@@ -1,5 +1,3 @@
|
||||
import type { ToolSet } from 'ai';
|
||||
|
||||
import {
|
||||
CopilotProviderSideError,
|
||||
metrics,
|
||||
@@ -11,6 +9,7 @@ import {
|
||||
type NativeLlmRequest,
|
||||
} from '../../../../native';
|
||||
import type { NodeTextMiddleware } from '../../config';
|
||||
import type { CopilotToolSet } from '../../tools';
|
||||
import { buildNativeRequest, NativeProviderAdapter } from '../native';
|
||||
import { CopilotProvider } from '../provider';
|
||||
import type {
|
||||
@@ -20,7 +19,11 @@ import type {
|
||||
StreamObject,
|
||||
} from '../types';
|
||||
import { CopilotProviderType, ModelOutputType } from '../types';
|
||||
import { getGoogleAuth, getVertexAnthropicBaseUrl } from '../utils';
|
||||
import {
|
||||
getGoogleAuth,
|
||||
getVertexAnthropicBaseUrl,
|
||||
type VertexAnthropicProviderConfig,
|
||||
} from '../utils';
|
||||
|
||||
export abstract class AnthropicProvider<T> extends CopilotProvider<T> {
|
||||
private handleError(e: any) {
|
||||
@@ -36,22 +39,16 @@ export abstract class AnthropicProvider<T> extends CopilotProvider<T> {
|
||||
|
||||
private async createNativeConfig(): Promise<NativeLlmBackendConfig> {
|
||||
if (this.type === CopilotProviderType.AnthropicVertex) {
|
||||
const auth = await getGoogleAuth(this.config as any, 'anthropic');
|
||||
const headers = auth.headers();
|
||||
const authorization =
|
||||
headers.Authorization ||
|
||||
(headers as Record<string, string | undefined>).authorization;
|
||||
const token =
|
||||
typeof authorization === 'string'
|
||||
? authorization.replace(/^Bearer\s+/i, '')
|
||||
: '';
|
||||
const baseUrl =
|
||||
getVertexAnthropicBaseUrl(this.config as any) || auth.baseUrl;
|
||||
const config = this.config as VertexAnthropicProviderConfig;
|
||||
const auth = await getGoogleAuth(config, 'anthropic');
|
||||
const { Authorization: authHeader } = auth.headers();
|
||||
const token = authHeader.replace(/^Bearer\s+/i, '');
|
||||
const baseUrl = getVertexAnthropicBaseUrl(config) || auth.baseUrl;
|
||||
return {
|
||||
base_url: baseUrl || '',
|
||||
auth_token: token,
|
||||
request_layer: 'vertex',
|
||||
headers,
|
||||
request_layer: 'vertex_anthropic',
|
||||
headers: { Authorization: authHeader },
|
||||
};
|
||||
}
|
||||
|
||||
@@ -65,7 +62,7 @@ export abstract class AnthropicProvider<T> extends CopilotProvider<T> {
|
||||
|
||||
private createAdapter(
|
||||
backendConfig: NativeLlmBackendConfig,
|
||||
tools: ToolSet,
|
||||
tools: CopilotToolSet,
|
||||
nodeTextMiddleware?: NodeTextMiddleware[]
|
||||
) {
|
||||
return new NativeProviderAdapter(
|
||||
@@ -93,8 +90,12 @@ export abstract class AnthropicProvider<T> extends CopilotProvider<T> {
|
||||
options: CopilotChatOptions = {}
|
||||
): Promise<string> {
|
||||
const fullCond = { ...cond, outputType: ModelOutputType.Text };
|
||||
await this.checkParams({ cond: fullCond, messages, options });
|
||||
const model = this.selectModel(fullCond);
|
||||
const normalizedCond = await this.checkParams({
|
||||
cond: fullCond,
|
||||
messages,
|
||||
options,
|
||||
});
|
||||
const model = this.selectModel(normalizedCond);
|
||||
|
||||
try {
|
||||
metrics.ai.counter('chat_text_calls').add(1, this.metricLabels(model.id));
|
||||
@@ -102,11 +103,13 @@ export abstract class AnthropicProvider<T> extends CopilotProvider<T> {
|
||||
const tools = await this.getTools(options, model.id);
|
||||
const middleware = this.getActiveProviderMiddleware();
|
||||
const reasoning = this.getReasoning(options, model.id);
|
||||
const cap = this.getAttachCapability(model, ModelOutputType.Text);
|
||||
const { request } = await buildNativeRequest({
|
||||
model: model.id,
|
||||
messages,
|
||||
options,
|
||||
tools,
|
||||
attachmentCapability: cap,
|
||||
reasoning,
|
||||
middleware,
|
||||
});
|
||||
@@ -115,7 +118,7 @@ export abstract class AnthropicProvider<T> extends CopilotProvider<T> {
|
||||
tools,
|
||||
middleware.node?.text
|
||||
);
|
||||
return await adapter.text(request, options.signal);
|
||||
return await adapter.text(request, options.signal, messages);
|
||||
} catch (e: any) {
|
||||
metrics.ai
|
||||
.counter('chat_text_errors')
|
||||
@@ -130,8 +133,12 @@ export abstract class AnthropicProvider<T> extends CopilotProvider<T> {
|
||||
options: CopilotChatOptions = {}
|
||||
): AsyncIterable<string> {
|
||||
const fullCond = { ...cond, outputType: ModelOutputType.Text };
|
||||
await this.checkParams({ cond: fullCond, messages, options });
|
||||
const model = this.selectModel(fullCond);
|
||||
const normalizedCond = await this.checkParams({
|
||||
cond: fullCond,
|
||||
messages,
|
||||
options,
|
||||
});
|
||||
const model = this.selectModel(normalizedCond);
|
||||
|
||||
try {
|
||||
metrics.ai
|
||||
@@ -140,11 +147,13 @@ export abstract class AnthropicProvider<T> extends CopilotProvider<T> {
|
||||
const backendConfig = await this.createNativeConfig();
|
||||
const tools = await this.getTools(options, model.id);
|
||||
const middleware = this.getActiveProviderMiddleware();
|
||||
const cap = this.getAttachCapability(model, ModelOutputType.Text);
|
||||
const { request } = await buildNativeRequest({
|
||||
model: model.id,
|
||||
messages,
|
||||
options,
|
||||
tools,
|
||||
attachmentCapability: cap,
|
||||
reasoning: this.getReasoning(options, model.id),
|
||||
middleware,
|
||||
});
|
||||
@@ -153,7 +162,11 @@ export abstract class AnthropicProvider<T> extends CopilotProvider<T> {
|
||||
tools,
|
||||
middleware.node?.text
|
||||
);
|
||||
for await (const chunk of adapter.streamText(request, options.signal)) {
|
||||
for await (const chunk of adapter.streamText(
|
||||
request,
|
||||
options.signal,
|
||||
messages
|
||||
)) {
|
||||
yield chunk;
|
||||
}
|
||||
} catch (e: any) {
|
||||
@@ -170,8 +183,12 @@ export abstract class AnthropicProvider<T> extends CopilotProvider<T> {
|
||||
options: CopilotChatOptions = {}
|
||||
): AsyncIterable<StreamObject> {
|
||||
const fullCond = { ...cond, outputType: ModelOutputType.Object };
|
||||
await this.checkParams({ cond: fullCond, messages, options });
|
||||
const model = this.selectModel(fullCond);
|
||||
const normalizedCond = await this.checkParams({
|
||||
cond: fullCond,
|
||||
messages,
|
||||
options,
|
||||
});
|
||||
const model = this.selectModel(normalizedCond);
|
||||
|
||||
try {
|
||||
metrics.ai
|
||||
@@ -180,11 +197,13 @@ export abstract class AnthropicProvider<T> extends CopilotProvider<T> {
|
||||
const backendConfig = await this.createNativeConfig();
|
||||
const tools = await this.getTools(options, model.id);
|
||||
const middleware = this.getActiveProviderMiddleware();
|
||||
const cap = this.getAttachCapability(model, ModelOutputType.Object);
|
||||
const { request } = await buildNativeRequest({
|
||||
model: model.id,
|
||||
messages,
|
||||
options,
|
||||
tools,
|
||||
attachmentCapability: cap,
|
||||
reasoning: this.getReasoning(options, model.id),
|
||||
middleware,
|
||||
});
|
||||
@@ -193,7 +212,11 @@ export abstract class AnthropicProvider<T> extends CopilotProvider<T> {
|
||||
tools,
|
||||
middleware.node?.text
|
||||
);
|
||||
for await (const chunk of adapter.streamObject(request, options.signal)) {
|
||||
for await (const chunk of adapter.streamObject(
|
||||
request,
|
||||
options.signal,
|
||||
messages
|
||||
)) {
|
||||
yield chunk;
|
||||
}
|
||||
} catch (e: any) {
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import z from 'zod';
|
||||
|
||||
import { IMAGE_ATTACHMENT_CAPABILITY } from '../attachments';
|
||||
import { CopilotProviderType, ModelInputType, ModelOutputType } from '../types';
|
||||
import { AnthropicProvider } from './anthropic';
|
||||
|
||||
@@ -23,6 +24,7 @@ export class AnthropicOfficialProvider extends AnthropicProvider<AnthropicOffici
|
||||
{
|
||||
input: [ModelInputType.Text, ModelInputType.Image],
|
||||
output: [ModelOutputType.Text, ModelOutputType.Object],
|
||||
attachments: IMAGE_ATTACHMENT_CAPABILITY,
|
||||
},
|
||||
],
|
||||
},
|
||||
@@ -33,6 +35,7 @@ export class AnthropicOfficialProvider extends AnthropicProvider<AnthropicOffici
|
||||
{
|
||||
input: [ModelInputType.Text, ModelInputType.Image],
|
||||
output: [ModelOutputType.Text, ModelOutputType.Object],
|
||||
attachments: IMAGE_ATTACHMENT_CAPABILITY,
|
||||
},
|
||||
],
|
||||
},
|
||||
@@ -43,6 +46,7 @@ export class AnthropicOfficialProvider extends AnthropicProvider<AnthropicOffici
|
||||
{
|
||||
input: [ModelInputType.Text, ModelInputType.Image],
|
||||
output: [ModelOutputType.Text, ModelOutputType.Object],
|
||||
attachments: IMAGE_ATTACHMENT_CAPABILITY,
|
||||
},
|
||||
],
|
||||
},
|
||||
|
||||
@@ -1,18 +1,14 @@
|
||||
import {
|
||||
createVertexAnthropic,
|
||||
type GoogleVertexAnthropicProvider,
|
||||
type GoogleVertexAnthropicProviderSettings,
|
||||
} from '@ai-sdk/google-vertex/anthropic';
|
||||
|
||||
import { IMAGE_ATTACHMENT_CAPABILITY } from '../attachments';
|
||||
import { CopilotProviderType, ModelInputType, ModelOutputType } from '../types';
|
||||
import {
|
||||
getGoogleAuth,
|
||||
getVertexAnthropicBaseUrl,
|
||||
VertexModelListSchema,
|
||||
type VertexProviderConfig,
|
||||
} from '../utils';
|
||||
import { AnthropicProvider } from './anthropic';
|
||||
|
||||
export type AnthropicVertexConfig = GoogleVertexAnthropicProviderSettings;
|
||||
export type AnthropicVertexConfig = VertexProviderConfig;
|
||||
|
||||
export class AnthropicVertexProvider extends AnthropicProvider<AnthropicVertexConfig> {
|
||||
override readonly type = CopilotProviderType.AnthropicVertex;
|
||||
@@ -25,6 +21,7 @@ export class AnthropicVertexProvider extends AnthropicProvider<AnthropicVertexCo
|
||||
{
|
||||
input: [ModelInputType.Text, ModelInputType.Image],
|
||||
output: [ModelOutputType.Text, ModelOutputType.Object],
|
||||
attachments: IMAGE_ATTACHMENT_CAPABILITY,
|
||||
},
|
||||
],
|
||||
},
|
||||
@@ -35,6 +32,7 @@ export class AnthropicVertexProvider extends AnthropicProvider<AnthropicVertexCo
|
||||
{
|
||||
input: [ModelInputType.Text, ModelInputType.Image],
|
||||
output: [ModelOutputType.Text, ModelOutputType.Object],
|
||||
attachments: IMAGE_ATTACHMENT_CAPABILITY,
|
||||
},
|
||||
],
|
||||
},
|
||||
@@ -45,23 +43,17 @@ export class AnthropicVertexProvider extends AnthropicProvider<AnthropicVertexCo
|
||||
{
|
||||
input: [ModelInputType.Text, ModelInputType.Image],
|
||||
output: [ModelOutputType.Text, ModelOutputType.Object],
|
||||
attachments: IMAGE_ATTACHMENT_CAPABILITY,
|
||||
},
|
||||
],
|
||||
},
|
||||
];
|
||||
|
||||
protected instance!: GoogleVertexAnthropicProvider;
|
||||
|
||||
override configured(): boolean {
|
||||
if (!this.config.location || !this.config.googleAuthOptions) return false;
|
||||
return !!this.config.project || !!getVertexAnthropicBaseUrl(this.config);
|
||||
}
|
||||
|
||||
override setup() {
|
||||
super.setup();
|
||||
this.instance = createVertexAnthropic(this.config);
|
||||
}
|
||||
|
||||
override async refreshOnlineModels() {
|
||||
try {
|
||||
const { baseUrl, headers } = await getGoogleAuth(
|
||||
|
||||
@@ -0,0 +1,233 @@
|
||||
import type {
|
||||
ModelAttachmentCapability,
|
||||
PromptAttachment,
|
||||
PromptAttachmentKind,
|
||||
PromptAttachmentSourceKind,
|
||||
PromptMessage,
|
||||
} from './types';
|
||||
import { inferMimeType } from './utils';
|
||||
|
||||
export const IMAGE_ATTACHMENT_CAPABILITY: ModelAttachmentCapability = {
|
||||
kinds: ['image'],
|
||||
sourceKinds: ['url', 'data'],
|
||||
allowRemoteUrls: true,
|
||||
};
|
||||
|
||||
export const GEMINI_ATTACHMENT_CAPABILITY: ModelAttachmentCapability = {
|
||||
kinds: ['image', 'audio', 'file'],
|
||||
sourceKinds: ['url', 'data', 'bytes', 'file_handle'],
|
||||
allowRemoteUrls: true,
|
||||
};
|
||||
|
||||
export type CanonicalPromptAttachment = {
|
||||
kind: PromptAttachmentKind;
|
||||
sourceKind: PromptAttachmentSourceKind;
|
||||
mediaType?: string;
|
||||
source: Record<string, unknown>;
|
||||
isRemote: boolean;
|
||||
};
|
||||
|
||||
function parseDataUrl(url: string) {
|
||||
if (!url.startsWith('data:')) {
|
||||
return null;
|
||||
}
|
||||
|
||||
const commaIndex = url.indexOf(',');
|
||||
if (commaIndex === -1) {
|
||||
return null;
|
||||
}
|
||||
|
||||
const meta = url.slice(5, commaIndex);
|
||||
const payload = url.slice(commaIndex + 1);
|
||||
const parts = meta.split(';');
|
||||
const mediaType = parts[0] || 'text/plain;charset=US-ASCII';
|
||||
const isBase64 = parts.includes('base64');
|
||||
|
||||
return {
|
||||
mediaType,
|
||||
data: isBase64
|
||||
? payload
|
||||
: Buffer.from(decodeURIComponent(payload), 'utf8').toString('base64'),
|
||||
};
|
||||
}
|
||||
|
||||
function attachmentTypeFromMediaType(mediaType: string): PromptAttachmentKind {
|
||||
if (mediaType.startsWith('image/')) {
|
||||
return 'image';
|
||||
}
|
||||
if (mediaType.startsWith('audio/')) {
|
||||
return 'audio';
|
||||
}
|
||||
return 'file';
|
||||
}
|
||||
|
||||
function attachmentKindFromHintOrMediaType(
|
||||
hint: PromptAttachmentKind | undefined,
|
||||
mediaType: string | undefined
|
||||
): PromptAttachmentKind {
|
||||
if (hint) return hint;
|
||||
return attachmentTypeFromMediaType(mediaType || '');
|
||||
}
|
||||
|
||||
function toBase64Data(data: string, encoding: 'base64' | 'utf8' = 'base64') {
|
||||
return encoding === 'base64'
|
||||
? data
|
||||
: Buffer.from(data, 'utf8').toString('base64');
|
||||
}
|
||||
|
||||
function appendAttachMetadata(
|
||||
source: Record<string, unknown>,
|
||||
attachment: Exclude<PromptAttachment, string> & Record<string, unknown>
|
||||
) {
|
||||
if (attachment.fileName) {
|
||||
source.file_name = attachment.fileName;
|
||||
}
|
||||
if (attachment.providerHint) {
|
||||
source.provider_hint = attachment.providerHint;
|
||||
}
|
||||
return source;
|
||||
}
|
||||
|
||||
export function promptAttachmentHasSource(
|
||||
attachment: PromptAttachment
|
||||
): boolean {
|
||||
if (typeof attachment === 'string') {
|
||||
return !!attachment.trim();
|
||||
}
|
||||
|
||||
if ('attachment' in attachment) {
|
||||
return !!attachment.attachment;
|
||||
}
|
||||
|
||||
switch (attachment.kind) {
|
||||
case 'url':
|
||||
return !!attachment.url;
|
||||
case 'data':
|
||||
case 'bytes':
|
||||
return !!attachment.data;
|
||||
case 'file_handle':
|
||||
return !!attachment.fileHandle;
|
||||
}
|
||||
}
|
||||
|
||||
export async function canonicalizePromptAttachment(
|
||||
attachment: PromptAttachment,
|
||||
message: Pick<PromptMessage, 'params'>
|
||||
): Promise<CanonicalPromptAttachment> {
|
||||
const fallbackMimeType =
|
||||
typeof message.params?.mimetype === 'string'
|
||||
? message.params.mimetype
|
||||
: undefined;
|
||||
|
||||
if (typeof attachment === 'string') {
|
||||
const dataUrl = parseDataUrl(attachment);
|
||||
const mediaType =
|
||||
fallbackMimeType ??
|
||||
dataUrl?.mediaType ??
|
||||
(await inferMimeType(attachment));
|
||||
const kind = attachmentKindFromHintOrMediaType(undefined, mediaType);
|
||||
if (dataUrl) {
|
||||
return {
|
||||
kind,
|
||||
sourceKind: 'data',
|
||||
mediaType,
|
||||
isRemote: false,
|
||||
source: {
|
||||
media_type: mediaType || dataUrl.mediaType,
|
||||
data: dataUrl.data,
|
||||
},
|
||||
};
|
||||
}
|
||||
|
||||
return {
|
||||
kind,
|
||||
sourceKind: 'url',
|
||||
mediaType,
|
||||
isRemote: /^https?:\/\//.test(attachment),
|
||||
source: { url: attachment, media_type: mediaType },
|
||||
};
|
||||
}
|
||||
|
||||
if ('attachment' in attachment) {
|
||||
return await canonicalizePromptAttachment(
|
||||
{
|
||||
kind: 'url',
|
||||
url: attachment.attachment,
|
||||
mimeType: attachment.mimeType,
|
||||
},
|
||||
message
|
||||
);
|
||||
}
|
||||
|
||||
if (attachment.kind === 'url') {
|
||||
const dataUrl = parseDataUrl(attachment.url);
|
||||
const mediaType =
|
||||
attachment.mimeType ??
|
||||
fallbackMimeType ??
|
||||
dataUrl?.mediaType ??
|
||||
(await inferMimeType(attachment.url));
|
||||
const kind = attachmentKindFromHintOrMediaType(
|
||||
attachment.providerHint?.kind,
|
||||
mediaType
|
||||
);
|
||||
if (dataUrl) {
|
||||
return {
|
||||
kind,
|
||||
sourceKind: 'data',
|
||||
mediaType,
|
||||
isRemote: false,
|
||||
source: appendAttachMetadata(
|
||||
{ media_type: mediaType || dataUrl.mediaType, data: dataUrl.data },
|
||||
attachment
|
||||
),
|
||||
};
|
||||
}
|
||||
|
||||
return {
|
||||
kind,
|
||||
sourceKind: 'url',
|
||||
mediaType,
|
||||
isRemote: /^https?:\/\//.test(attachment.url),
|
||||
source: appendAttachMetadata(
|
||||
{ url: attachment.url, media_type: mediaType },
|
||||
attachment
|
||||
),
|
||||
};
|
||||
}
|
||||
|
||||
if (attachment.kind === 'data' || attachment.kind === 'bytes') {
|
||||
return {
|
||||
kind: attachmentKindFromHintOrMediaType(
|
||||
attachment.providerHint?.kind,
|
||||
attachment.mimeType
|
||||
),
|
||||
sourceKind: attachment.kind,
|
||||
mediaType: attachment.mimeType,
|
||||
isRemote: false,
|
||||
source: appendAttachMetadata(
|
||||
{
|
||||
media_type: attachment.mimeType,
|
||||
data: toBase64Data(
|
||||
attachment.data,
|
||||
attachment.kind === 'data' ? attachment.encoding : 'base64'
|
||||
),
|
||||
},
|
||||
attachment
|
||||
),
|
||||
};
|
||||
}
|
||||
|
||||
return {
|
||||
kind: attachmentKindFromHintOrMediaType(
|
||||
attachment.providerHint?.kind,
|
||||
attachment.mimeType
|
||||
),
|
||||
sourceKind: 'file_handle',
|
||||
mediaType: attachment.mimeType,
|
||||
isRemote: false,
|
||||
source: appendAttachMetadata(
|
||||
{ file_handle: attachment.fileHandle, media_type: attachment.mimeType },
|
||||
attachment
|
||||
),
|
||||
};
|
||||
}
|
||||
@@ -19,6 +19,7 @@ import type {
|
||||
PromptMessage,
|
||||
} from './types';
|
||||
import { CopilotProviderType, ModelInputType, ModelOutputType } from './types';
|
||||
import { promptAttachmentMimeType, promptAttachmentToUrl } from './utils';
|
||||
|
||||
export type FalConfig = {
|
||||
apiKey: string;
|
||||
@@ -183,13 +184,14 @@ export class FalProvider extends CopilotProvider<FalConfig> {
|
||||
return {
|
||||
model_name: options.modelName || undefined,
|
||||
image_url: attachments
|
||||
?.map(v =>
|
||||
typeof v === 'string'
|
||||
? v
|
||||
: v.mimeType.startsWith('image/')
|
||||
? v.attachment
|
||||
: undefined
|
||||
)
|
||||
?.map(v => {
|
||||
const url = promptAttachmentToUrl(v);
|
||||
const mediaType = promptAttachmentMimeType(
|
||||
v,
|
||||
typeof params?.mimetype === 'string' ? params.mimetype : undefined
|
||||
);
|
||||
return url && mediaType?.startsWith('image/') ? url : undefined;
|
||||
})
|
||||
.find(v => !!v),
|
||||
prompt: content.trim(),
|
||||
loras: lora.length ? lora : undefined,
|
||||
|
||||
@@ -1,87 +1,94 @@
|
||||
import type {
|
||||
GoogleGenerativeAIProvider,
|
||||
GoogleGenerativeAIProviderOptions,
|
||||
} from '@ai-sdk/google';
|
||||
import type { GoogleVertexProvider } from '@ai-sdk/google-vertex';
|
||||
import {
|
||||
AISDKError,
|
||||
type EmbeddingModel,
|
||||
embedMany,
|
||||
generateObject,
|
||||
generateText,
|
||||
JSONParseError,
|
||||
stepCountIs,
|
||||
streamText,
|
||||
} from 'ai';
|
||||
import { setTimeout as delay } from 'node:timers/promises';
|
||||
|
||||
import { ZodError } from 'zod';
|
||||
|
||||
import {
|
||||
CopilotPromptInvalid,
|
||||
CopilotProviderSideError,
|
||||
metrics,
|
||||
OneMB,
|
||||
readResponseBufferWithLimit,
|
||||
safeFetch,
|
||||
UserFriendlyError,
|
||||
} from '../../../../base';
|
||||
import { sniffMime } from '../../../../base/storage/providers/utils';
|
||||
import {
|
||||
llmDispatchStream,
|
||||
llmEmbeddingDispatch,
|
||||
llmStructuredDispatch,
|
||||
type NativeLlmBackendConfig,
|
||||
type NativeLlmEmbeddingRequest,
|
||||
type NativeLlmRequest,
|
||||
type NativeLlmStructuredRequest,
|
||||
} from '../../../../native';
|
||||
import type { NodeTextMiddleware } from '../../config';
|
||||
import type { CopilotToolSet } from '../../tools';
|
||||
import {
|
||||
buildNativeEmbeddingRequest,
|
||||
buildNativeRequest,
|
||||
buildNativeStructuredRequest,
|
||||
NativeProviderAdapter,
|
||||
parseNativeStructuredOutput,
|
||||
StructuredResponseParseError,
|
||||
} from '../native';
|
||||
import { CopilotProvider } from '../provider';
|
||||
import type {
|
||||
CopilotChatOptions,
|
||||
CopilotEmbeddingOptions,
|
||||
CopilotImageOptions,
|
||||
CopilotProviderModel,
|
||||
CopilotStructuredOptions,
|
||||
ModelConditions,
|
||||
PromptAttachment,
|
||||
PromptMessage,
|
||||
StreamObject,
|
||||
} from '../types';
|
||||
import { ModelOutputType } from '../types';
|
||||
import {
|
||||
chatToGPTMessage,
|
||||
StreamObjectParser,
|
||||
TextStreamParser,
|
||||
} from '../utils';
|
||||
import { promptAttachmentMimeType, promptAttachmentToUrl } from '../utils';
|
||||
|
||||
export const DEFAULT_DIMENSIONS = 256;
|
||||
const GEMINI_REMOTE_ATTACHMENT_MAX_BYTES = 64 * OneMB;
|
||||
const TRUSTED_ATTACHMENT_HOST_SUFFIXES = ['cdn.affine.pro'];
|
||||
const GEMINI_RETRY_INITIAL_DELAY_MS = 2_000;
|
||||
|
||||
function normalizeMimeType(mediaType?: string) {
|
||||
return mediaType?.split(';', 1)[0]?.trim() || 'application/octet-stream';
|
||||
}
|
||||
|
||||
function isYoutubeUrl(url: URL) {
|
||||
const hostname = url.hostname.toLowerCase();
|
||||
if (hostname === 'youtu.be') {
|
||||
return /^\/[\w-]+$/.test(url.pathname);
|
||||
}
|
||||
|
||||
if (hostname !== 'youtube.com' && hostname !== 'www.youtube.com') {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (url.pathname !== '/watch') {
|
||||
return false;
|
||||
}
|
||||
|
||||
return !!url.searchParams.get('v');
|
||||
}
|
||||
|
||||
function isGeminiFileUrl(url: URL, baseUrl: string) {
|
||||
try {
|
||||
const base = new URL(baseUrl);
|
||||
const basePath = base.pathname.replace(/\/+$/, '');
|
||||
return (
|
||||
url.origin === base.origin &&
|
||||
url.pathname.startsWith(`${basePath}/files/`)
|
||||
);
|
||||
} catch {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
export abstract class GeminiProvider<T> extends CopilotProvider<T> {
|
||||
protected abstract instance:
|
||||
| GoogleGenerativeAIProvider
|
||||
| GoogleVertexProvider;
|
||||
|
||||
private getThinkingConfig(
|
||||
model: string,
|
||||
options: { includeThoughts: boolean; useDynamicBudget?: boolean }
|
||||
): NonNullable<GoogleGenerativeAIProviderOptions['thinkingConfig']> {
|
||||
if (this.isGemini3Model(model)) {
|
||||
return {
|
||||
includeThoughts: options.includeThoughts,
|
||||
thinkingLevel: 'high',
|
||||
};
|
||||
}
|
||||
|
||||
return {
|
||||
includeThoughts: options.includeThoughts,
|
||||
thinkingBudget: options.useDynamicBudget ? -1 : 12000,
|
||||
};
|
||||
}
|
||||
|
||||
private getEmbeddingModel(model: string) {
|
||||
const provider = this.instance as typeof this.instance & {
|
||||
embeddingModel?: (modelId: string) => EmbeddingModel;
|
||||
textEmbeddingModel?: (modelId: string) => EmbeddingModel;
|
||||
};
|
||||
|
||||
return (
|
||||
provider.embeddingModel?.(model) ?? provider.textEmbeddingModel?.(model)
|
||||
);
|
||||
}
|
||||
protected abstract createNativeConfig(): Promise<NativeLlmBackendConfig>;
|
||||
|
||||
private handleError(e: any) {
|
||||
if (e instanceof UserFriendlyError) {
|
||||
return e;
|
||||
} else if (e instanceof AISDKError) {
|
||||
this.logger.error('Throw error from ai sdk:', e);
|
||||
return new CopilotProviderSideError({
|
||||
provider: this.type,
|
||||
kind: e.name || 'unknown',
|
||||
message: e.message,
|
||||
});
|
||||
} else {
|
||||
return new CopilotProviderSideError({
|
||||
provider: this.type,
|
||||
@@ -91,37 +98,261 @@ export abstract class GeminiProvider<T> extends CopilotProvider<T> {
|
||||
}
|
||||
}
|
||||
|
||||
protected createNativeDispatch(backendConfig: NativeLlmBackendConfig) {
|
||||
return (request: NativeLlmRequest, signal?: AbortSignal) =>
|
||||
llmDispatchStream('gemini', backendConfig, request, signal);
|
||||
}
|
||||
|
||||
protected createNativeStructuredDispatch(
|
||||
backendConfig: NativeLlmBackendConfig
|
||||
) {
|
||||
return (request: NativeLlmStructuredRequest) =>
|
||||
llmStructuredDispatch('gemini', backendConfig, request);
|
||||
}
|
||||
|
||||
protected createNativeEmbeddingDispatch(
|
||||
backendConfig: NativeLlmBackendConfig
|
||||
) {
|
||||
return (request: NativeLlmEmbeddingRequest) =>
|
||||
llmEmbeddingDispatch('gemini', backendConfig, request);
|
||||
}
|
||||
|
||||
protected createNativeAdapter(
|
||||
backendConfig: NativeLlmBackendConfig,
|
||||
tools: CopilotToolSet,
|
||||
nodeTextMiddleware?: NodeTextMiddleware[]
|
||||
) {
|
||||
return new NativeProviderAdapter(
|
||||
this.createNativeDispatch(backendConfig),
|
||||
tools,
|
||||
this.MAX_STEPS,
|
||||
{ nodeTextMiddleware }
|
||||
);
|
||||
}
|
||||
|
||||
protected async fetchRemoteAttach(url: string, signal?: AbortSignal) {
|
||||
const parsed = new URL(url);
|
||||
const response = await safeFetch(
|
||||
parsed,
|
||||
{ method: 'GET', signal },
|
||||
this.buildAttachFetchOptions(parsed)
|
||||
);
|
||||
if (!response.ok) {
|
||||
throw new Error(
|
||||
`Failed to fetch attachment: ${response.status} ${response.statusText}`
|
||||
);
|
||||
}
|
||||
const buffer = await readResponseBufferWithLimit(
|
||||
response,
|
||||
GEMINI_REMOTE_ATTACHMENT_MAX_BYTES
|
||||
);
|
||||
const headerMimeType = normalizeMimeType(
|
||||
response.headers.get('content-type') || ''
|
||||
);
|
||||
return {
|
||||
data: buffer.toString('base64'),
|
||||
mimeType: normalizeMimeType(sniffMime(buffer, headerMimeType)),
|
||||
};
|
||||
}
|
||||
|
||||
private buildAttachFetchOptions(url: URL) {
|
||||
const baseOptions = { timeoutMs: 15_000, maxRedirects: 3 } as const;
|
||||
if (!env.prod) {
|
||||
return { ...baseOptions, allowPrivateOrigins: new Set([url.origin]) };
|
||||
}
|
||||
|
||||
const trustedOrigins = new Set<string>();
|
||||
const protocol = this.AFFiNEConfig.server.https ? 'https:' : 'http:';
|
||||
const port = this.AFFiNEConfig.server.port;
|
||||
const isDefaultPort =
|
||||
(protocol === 'https:' && port === 443) ||
|
||||
(protocol === 'http:' && port === 80);
|
||||
|
||||
const addHostOrigin = (host: string) => {
|
||||
if (!host) return;
|
||||
try {
|
||||
const parsed = new URL(`${protocol}//${host}`);
|
||||
if (!parsed.port && !isDefaultPort) {
|
||||
parsed.port = String(port);
|
||||
}
|
||||
trustedOrigins.add(parsed.origin);
|
||||
} catch {
|
||||
// ignore invalid host config entries
|
||||
}
|
||||
};
|
||||
|
||||
if (this.AFFiNEConfig.server.externalUrl) {
|
||||
try {
|
||||
trustedOrigins.add(
|
||||
new URL(this.AFFiNEConfig.server.externalUrl).origin
|
||||
);
|
||||
} catch {
|
||||
// ignore invalid external URL
|
||||
}
|
||||
}
|
||||
|
||||
addHostOrigin(this.AFFiNEConfig.server.host);
|
||||
for (const host of this.AFFiNEConfig.server.hosts) {
|
||||
addHostOrigin(host);
|
||||
}
|
||||
|
||||
const hostname = url.hostname.toLowerCase();
|
||||
const trustedByHost = TRUSTED_ATTACHMENT_HOST_SUFFIXES.some(
|
||||
suffix => hostname === suffix || hostname.endsWith(`.${suffix}`)
|
||||
);
|
||||
if (trustedOrigins.has(url.origin) || trustedByHost) {
|
||||
return { ...baseOptions, allowPrivateOrigins: new Set([url.origin]) };
|
||||
}
|
||||
|
||||
return baseOptions;
|
||||
}
|
||||
|
||||
private shouldInlineRemoteAttach(url: URL, config: NativeLlmBackendConfig) {
|
||||
switch (config.request_layer) {
|
||||
case 'gemini_api':
|
||||
if (url.protocol !== 'http:' && url.protocol !== 'https:') return false;
|
||||
return !(isGeminiFileUrl(url, config.base_url) || isYoutubeUrl(url));
|
||||
case 'gemini_vertex':
|
||||
return false;
|
||||
default:
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
private toInlineAttach(
|
||||
attachment: PromptAttachment,
|
||||
mimeType: string,
|
||||
data: string
|
||||
): PromptAttachment {
|
||||
if (typeof attachment === 'string' || !('kind' in attachment)) {
|
||||
return { kind: 'bytes', data, mimeType };
|
||||
}
|
||||
|
||||
if (attachment.kind !== 'url') {
|
||||
return attachment;
|
||||
}
|
||||
|
||||
return {
|
||||
kind: 'bytes',
|
||||
data,
|
||||
mimeType,
|
||||
fileName: attachment.fileName,
|
||||
providerHint: attachment.providerHint,
|
||||
};
|
||||
}
|
||||
|
||||
protected async prepareMessages(
|
||||
messages: PromptMessage[],
|
||||
backendConfig: NativeLlmBackendConfig,
|
||||
signal?: AbortSignal
|
||||
): Promise<PromptMessage[]> {
|
||||
const prepared: PromptMessage[] = [];
|
||||
|
||||
for (const message of messages) {
|
||||
signal?.throwIfAborted();
|
||||
if (!Array.isArray(message.attachments) || !message.attachments.length) {
|
||||
prepared.push(message);
|
||||
continue;
|
||||
}
|
||||
|
||||
const attachments: PromptAttachment[] = [];
|
||||
let changed = false;
|
||||
for (const attachment of message.attachments) {
|
||||
signal?.throwIfAborted();
|
||||
const rawUrl = promptAttachmentToUrl(attachment);
|
||||
if (!rawUrl || rawUrl.startsWith('data:')) {
|
||||
attachments.push(attachment);
|
||||
continue;
|
||||
}
|
||||
|
||||
let parsed: URL;
|
||||
try {
|
||||
parsed = new URL(rawUrl);
|
||||
} catch {
|
||||
attachments.push(attachment);
|
||||
continue;
|
||||
}
|
||||
|
||||
if (!this.shouldInlineRemoteAttach(parsed, backendConfig)) {
|
||||
attachments.push(attachment);
|
||||
continue;
|
||||
}
|
||||
|
||||
const declaredMimeType = promptAttachmentMimeType(
|
||||
attachment,
|
||||
typeof message.params?.mimetype === 'string'
|
||||
? message.params.mimetype
|
||||
: undefined
|
||||
);
|
||||
const downloaded = await this.fetchRemoteAttach(rawUrl, signal);
|
||||
attachments.push(
|
||||
this.toInlineAttach(
|
||||
attachment,
|
||||
declaredMimeType
|
||||
? normalizeMimeType(declaredMimeType)
|
||||
: downloaded.mimeType,
|
||||
downloaded.data
|
||||
)
|
||||
);
|
||||
changed = true;
|
||||
}
|
||||
|
||||
prepared.push(changed ? { ...message, attachments } : message);
|
||||
}
|
||||
|
||||
return prepared;
|
||||
}
|
||||
|
||||
protected async waitForStructuredRetry(
|
||||
delayMs: number,
|
||||
signal?: AbortSignal
|
||||
) {
|
||||
await delay(delayMs, undefined, signal ? { signal } : undefined);
|
||||
}
|
||||
|
||||
async text(
|
||||
cond: ModelConditions,
|
||||
messages: PromptMessage[],
|
||||
options: CopilotChatOptions = {}
|
||||
): Promise<string> {
|
||||
const fullCond = { ...cond, outputType: ModelOutputType.Text };
|
||||
await this.checkParams({ cond: fullCond, messages, options });
|
||||
const model = this.selectModel(fullCond);
|
||||
const normalizedCond = await this.checkParams({
|
||||
cond: fullCond,
|
||||
messages,
|
||||
options,
|
||||
});
|
||||
const model = this.selectModel(normalizedCond);
|
||||
|
||||
try {
|
||||
metrics.ai.counter('chat_text_calls').add(1, { model: model.id });
|
||||
|
||||
const [system, msgs] = await chatToGPTMessage(messages);
|
||||
|
||||
const modelInstance = this.instance(model.id);
|
||||
const { text } = await generateText({
|
||||
model: modelInstance,
|
||||
system,
|
||||
messages: msgs,
|
||||
abortSignal: options.signal,
|
||||
providerOptions: {
|
||||
google: this.getGeminiOptions(options, model.id),
|
||||
},
|
||||
tools: await this.getTools(options, model.id),
|
||||
stopWhen: stepCountIs(this.MAX_STEPS),
|
||||
metrics.ai.counter('chat_text_calls').add(1, this.metricLabels(model.id));
|
||||
const backendConfig = await this.createNativeConfig();
|
||||
const msg = await this.prepareMessages(
|
||||
messages,
|
||||
backendConfig,
|
||||
options.signal
|
||||
);
|
||||
const tools = await this.getTools(options, model.id);
|
||||
const middleware = this.getActiveProviderMiddleware();
|
||||
const cap = this.getAttachCapability(model, ModelOutputType.Text);
|
||||
const { request } = await buildNativeRequest({
|
||||
model: model.id,
|
||||
messages: msg,
|
||||
options,
|
||||
tools,
|
||||
attachmentCapability: cap,
|
||||
reasoning: this.getReasoning(options, model.id),
|
||||
middleware,
|
||||
});
|
||||
|
||||
if (!text) throw new Error('Failed to generate text');
|
||||
return text.trim();
|
||||
const adapter = this.createNativeAdapter(
|
||||
backendConfig,
|
||||
tools,
|
||||
middleware.node?.text
|
||||
);
|
||||
return await adapter.text(request, options.signal, messages);
|
||||
} catch (e: any) {
|
||||
metrics.ai.counter('chat_text_errors').add(1, { model: model.id });
|
||||
metrics.ai
|
||||
.counter('chat_text_errors')
|
||||
.add(1, this.metricLabels(model.id));
|
||||
throw this.handleError(e);
|
||||
}
|
||||
}
|
||||
@@ -129,55 +360,65 @@ export abstract class GeminiProvider<T> extends CopilotProvider<T> {
|
||||
override async structure(
|
||||
cond: ModelConditions,
|
||||
messages: PromptMessage[],
|
||||
options: CopilotChatOptions = {}
|
||||
options: CopilotStructuredOptions = {}
|
||||
): Promise<string> {
|
||||
const fullCond = { ...cond, outputType: ModelOutputType.Structured };
|
||||
await this.checkParams({ cond: fullCond, messages, options });
|
||||
const model = this.selectModel(fullCond);
|
||||
const normalizedCond = await this.checkParams({
|
||||
cond: fullCond,
|
||||
messages,
|
||||
options,
|
||||
});
|
||||
const model = this.selectModel(normalizedCond);
|
||||
|
||||
try {
|
||||
metrics.ai.counter('chat_text_calls').add(1, { model: model.id });
|
||||
|
||||
const [system, msgs, schema] = await chatToGPTMessage(messages);
|
||||
if (!schema) {
|
||||
throw new CopilotPromptInvalid('Schema is required');
|
||||
}
|
||||
|
||||
const modelInstance = this.instance(model.id);
|
||||
const { object } = await generateObject({
|
||||
model: modelInstance,
|
||||
system,
|
||||
messages: msgs,
|
||||
schema,
|
||||
providerOptions: {
|
||||
google: {
|
||||
thinkingConfig: this.getThinkingConfig(model.id, {
|
||||
includeThoughts: false,
|
||||
useDynamicBudget: true,
|
||||
}),
|
||||
},
|
||||
},
|
||||
abortSignal: options.signal,
|
||||
maxRetries: options.maxRetries || 3,
|
||||
experimental_repairText: async ({ text, error }) => {
|
||||
if (error instanceof JSONParseError) {
|
||||
// strange fixed response, temporarily replace it
|
||||
const ret = text.replaceAll(/^ny\n/g, ' ').trim();
|
||||
if (ret.startsWith('```') || ret.endsWith('```')) {
|
||||
return ret
|
||||
.replace(/```[\w\s]+\n/g, '')
|
||||
.replace(/\n```/g, '')
|
||||
.trim();
|
||||
}
|
||||
return ret;
|
||||
}
|
||||
return null;
|
||||
},
|
||||
metrics.ai.counter('chat_text_calls').add(1, this.metricLabels(model.id));
|
||||
const backendConfig = await this.createNativeConfig();
|
||||
const msg = await this.prepareMessages(
|
||||
messages,
|
||||
backendConfig,
|
||||
options.signal
|
||||
);
|
||||
const structuredDispatch =
|
||||
this.createNativeStructuredDispatch(backendConfig);
|
||||
const middleware = this.getActiveProviderMiddleware();
|
||||
const cap = this.getAttachCapability(model, ModelOutputType.Structured);
|
||||
const { request, schema } = await buildNativeStructuredRequest({
|
||||
model: model.id,
|
||||
messages: msg,
|
||||
options,
|
||||
attachmentCapability: cap,
|
||||
reasoning: this.getReasoning(options, model.id),
|
||||
responseSchema: options.schema,
|
||||
middleware,
|
||||
});
|
||||
|
||||
return JSON.stringify(object);
|
||||
const maxRetries = Math.max(options.maxRetries ?? 3, 0);
|
||||
for (let attempt = 0; ; attempt++) {
|
||||
try {
|
||||
const response = await structuredDispatch(request);
|
||||
const parsed = parseNativeStructuredOutput(response);
|
||||
const validated = schema.parse(parsed);
|
||||
return JSON.stringify(validated);
|
||||
} catch (error) {
|
||||
const isParsingError =
|
||||
error instanceof StructuredResponseParseError ||
|
||||
error instanceof ZodError;
|
||||
const retryableError =
|
||||
isParsingError || !(error instanceof UserFriendlyError);
|
||||
if (!retryableError || attempt >= maxRetries) {
|
||||
throw error;
|
||||
}
|
||||
if (!isParsingError) {
|
||||
await this.waitForStructuredRetry(
|
||||
GEMINI_RETRY_INITIAL_DELAY_MS * 2 ** attempt,
|
||||
options.signal
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
} catch (e: any) {
|
||||
metrics.ai.counter('chat_text_errors').add(1, { model: model.id });
|
||||
metrics.ai
|
||||
.counter('chat_text_errors')
|
||||
.add(1, this.metricLabels(model.id));
|
||||
throw this.handleError(e);
|
||||
}
|
||||
}
|
||||
@@ -188,29 +429,54 @@ export abstract class GeminiProvider<T> extends CopilotProvider<T> {
|
||||
options: CopilotChatOptions | CopilotImageOptions = {}
|
||||
): AsyncIterable<string> {
|
||||
const fullCond = { ...cond, outputType: ModelOutputType.Text };
|
||||
await this.checkParams({ cond: fullCond, messages, options });
|
||||
const model = this.selectModel(fullCond);
|
||||
const normalizedCond = await this.checkParams({
|
||||
cond: fullCond,
|
||||
messages,
|
||||
options,
|
||||
});
|
||||
const model = this.selectModel(normalizedCond);
|
||||
|
||||
try {
|
||||
metrics.ai.counter('chat_text_stream_calls').add(1, { model: model.id });
|
||||
const fullStream = await this.getFullStream(model, messages, options);
|
||||
const parser = new TextStreamParser();
|
||||
for await (const chunk of fullStream) {
|
||||
const result = parser.parse(chunk);
|
||||
yield result;
|
||||
if (options.signal?.aborted) {
|
||||
await fullStream.cancel();
|
||||
break;
|
||||
}
|
||||
}
|
||||
if (!options.signal?.aborted) {
|
||||
const footnotes = parser.end();
|
||||
if (footnotes.length) {
|
||||
yield `\n\n${footnotes}`;
|
||||
}
|
||||
metrics.ai
|
||||
.counter('chat_text_stream_calls')
|
||||
.add(1, this.metricLabels(model.id));
|
||||
const backendConfig = await this.createNativeConfig();
|
||||
const preparedMessages = await this.prepareMessages(
|
||||
messages,
|
||||
backendConfig,
|
||||
options.signal
|
||||
);
|
||||
const tools = await this.getTools(
|
||||
options as CopilotChatOptions,
|
||||
model.id
|
||||
);
|
||||
const middleware = this.getActiveProviderMiddleware();
|
||||
const cap = this.getAttachCapability(model, ModelOutputType.Text);
|
||||
const { request } = await buildNativeRequest({
|
||||
model: model.id,
|
||||
messages: preparedMessages,
|
||||
options: options as CopilotChatOptions,
|
||||
tools,
|
||||
attachmentCapability: cap,
|
||||
reasoning: this.getReasoning(options, model.id),
|
||||
middleware,
|
||||
});
|
||||
const adapter = this.createNativeAdapter(
|
||||
backendConfig,
|
||||
tools,
|
||||
middleware.node?.text
|
||||
);
|
||||
for await (const chunk of adapter.streamText(
|
||||
request,
|
||||
options.signal,
|
||||
messages
|
||||
)) {
|
||||
yield chunk;
|
||||
}
|
||||
} catch (e: any) {
|
||||
metrics.ai.counter('chat_text_stream_errors').add(1, { model: model.id });
|
||||
metrics.ai
|
||||
.counter('chat_text_stream_errors')
|
||||
.add(1, this.metricLabels(model.id));
|
||||
throw this.handleError(e);
|
||||
}
|
||||
}
|
||||
@@ -221,29 +487,51 @@ export abstract class GeminiProvider<T> extends CopilotProvider<T> {
|
||||
options: CopilotChatOptions = {}
|
||||
): AsyncIterable<StreamObject> {
|
||||
const fullCond = { ...cond, outputType: ModelOutputType.Object };
|
||||
await this.checkParams({ cond: fullCond, messages, options });
|
||||
const model = this.selectModel(fullCond);
|
||||
const normalizedCond = await this.checkParams({
|
||||
cond: fullCond,
|
||||
messages,
|
||||
options,
|
||||
});
|
||||
const model = this.selectModel(normalizedCond);
|
||||
|
||||
try {
|
||||
metrics.ai
|
||||
.counter('chat_object_stream_calls')
|
||||
.add(1, { model: model.id });
|
||||
const fullStream = await this.getFullStream(model, messages, options);
|
||||
const parser = new StreamObjectParser();
|
||||
for await (const chunk of fullStream) {
|
||||
const result = parser.parse(chunk);
|
||||
if (result) {
|
||||
yield result;
|
||||
}
|
||||
if (options.signal?.aborted) {
|
||||
await fullStream.cancel();
|
||||
break;
|
||||
}
|
||||
.add(1, this.metricLabels(model.id));
|
||||
const backendConfig = await this.createNativeConfig();
|
||||
const msg = await this.prepareMessages(
|
||||
messages,
|
||||
backendConfig,
|
||||
options.signal
|
||||
);
|
||||
const tools = await this.getTools(options, model.id);
|
||||
const middleware = this.getActiveProviderMiddleware();
|
||||
const cap = this.getAttachCapability(model, ModelOutputType.Object);
|
||||
const { request } = await buildNativeRequest({
|
||||
model: model.id,
|
||||
messages: msg,
|
||||
options,
|
||||
tools,
|
||||
attachmentCapability: cap,
|
||||
reasoning: this.getReasoning(options, model.id),
|
||||
middleware,
|
||||
});
|
||||
const adapter = this.createNativeAdapter(
|
||||
backendConfig,
|
||||
tools,
|
||||
middleware.node?.text
|
||||
);
|
||||
for await (const chunk of adapter.streamObject(
|
||||
request,
|
||||
options.signal,
|
||||
messages
|
||||
)) {
|
||||
yield chunk;
|
||||
}
|
||||
} catch (e: any) {
|
||||
metrics.ai
|
||||
.counter('chat_object_stream_errors')
|
||||
.add(1, { model: model.id });
|
||||
.add(1, this.metricLabels(model.id));
|
||||
throw this.handleError(e);
|
||||
}
|
||||
}
|
||||
@@ -253,76 +541,53 @@ export abstract class GeminiProvider<T> extends CopilotProvider<T> {
|
||||
messages: string | string[],
|
||||
options: CopilotEmbeddingOptions = { dimensions: DEFAULT_DIMENSIONS }
|
||||
): Promise<number[][]> {
|
||||
messages = Array.isArray(messages) ? messages : [messages];
|
||||
const values = Array.isArray(messages) ? messages : [messages];
|
||||
const fullCond = { ...cond, outputType: ModelOutputType.Embedding };
|
||||
await this.checkParams({ embeddings: messages, cond: fullCond, options });
|
||||
const model = this.selectModel(fullCond);
|
||||
const normalizedCond = await this.checkParams({
|
||||
embeddings: values,
|
||||
cond: fullCond,
|
||||
options,
|
||||
});
|
||||
const model = this.selectModel(normalizedCond);
|
||||
|
||||
try {
|
||||
metrics.ai
|
||||
.counter('generate_embedding_calls')
|
||||
.add(1, { model: model.id });
|
||||
|
||||
const modelInstance = this.getEmbeddingModel(model.id);
|
||||
if (!modelInstance) {
|
||||
throw new Error(`Embedding model is not available for ${model.id}`);
|
||||
}
|
||||
|
||||
const embeddings = await Promise.allSettled(
|
||||
messages.map(m =>
|
||||
embedMany({
|
||||
model: modelInstance,
|
||||
values: [m],
|
||||
maxRetries: 3,
|
||||
providerOptions: {
|
||||
google: {
|
||||
outputDimensionality: options.dimensions || DEFAULT_DIMENSIONS,
|
||||
taskType: 'RETRIEVAL_DOCUMENT',
|
||||
},
|
||||
},
|
||||
})
|
||||
)
|
||||
.add(1, this.metricLabels(model.id));
|
||||
const backendConfig = await this.createNativeConfig();
|
||||
const response = await this.createNativeEmbeddingDispatch(backendConfig)(
|
||||
buildNativeEmbeddingRequest({
|
||||
model: model.id,
|
||||
inputs: values,
|
||||
dimensions: options.dimensions || DEFAULT_DIMENSIONS,
|
||||
taskType: 'RETRIEVAL_DOCUMENT',
|
||||
})
|
||||
);
|
||||
|
||||
return embeddings
|
||||
.flatMap(e => (e.status === 'fulfilled' ? e.value.embeddings : null))
|
||||
.filter((v): v is number[] => !!v && Array.isArray(v));
|
||||
return response.embeddings;
|
||||
} catch (e: any) {
|
||||
metrics.ai
|
||||
.counter('generate_embedding_errors')
|
||||
.add(1, { model: model.id });
|
||||
.add(1, this.metricLabels(model.id));
|
||||
throw this.handleError(e);
|
||||
}
|
||||
}
|
||||
|
||||
private async getFullStream(
|
||||
model: CopilotProviderModel,
|
||||
messages: PromptMessage[],
|
||||
options: CopilotChatOptions = {}
|
||||
) {
|
||||
const [system, msgs] = await chatToGPTMessage(messages);
|
||||
const { fullStream } = streamText({
|
||||
model: this.instance(model.id),
|
||||
system,
|
||||
messages: msgs,
|
||||
abortSignal: options.signal,
|
||||
providerOptions: {
|
||||
google: this.getGeminiOptions(options, model.id),
|
||||
},
|
||||
tools: await this.getTools(options, model.id),
|
||||
stopWhen: stepCountIs(this.MAX_STEPS),
|
||||
});
|
||||
return fullStream;
|
||||
}
|
||||
|
||||
private getGeminiOptions(options: CopilotChatOptions, model: string) {
|
||||
const result: GoogleGenerativeAIProviderOptions = {};
|
||||
if (options?.reasoning && this.isReasoningModel(model)) {
|
||||
result.thinkingConfig = this.getThinkingConfig(model, {
|
||||
includeThoughts: true,
|
||||
});
|
||||
protected getReasoning(
|
||||
options: CopilotChatOptions | CopilotImageOptions,
|
||||
model: string
|
||||
): Record<string, unknown> | undefined {
|
||||
if (
|
||||
options &&
|
||||
'reasoning' in options &&
|
||||
options.reasoning &&
|
||||
this.isReasoningModel(model)
|
||||
) {
|
||||
return this.isGemini3Model(model)
|
||||
? { include_thoughts: true, thinking_level: 'high' }
|
||||
: { include_thoughts: true, thinking_budget: 12000 };
|
||||
}
|
||||
return result;
|
||||
|
||||
return undefined;
|
||||
}
|
||||
|
||||
private isGemini3Model(model: string) {
|
||||
|
||||
@@ -1,9 +1,7 @@
|
||||
import {
|
||||
createGoogleGenerativeAI,
|
||||
type GoogleGenerativeAIProvider,
|
||||
} from '@ai-sdk/google';
|
||||
import z from 'zod';
|
||||
|
||||
import type { NativeLlmBackendConfig } from '../../../../native';
|
||||
import { GEMINI_ATTACHMENT_CAPABILITY } from '../attachments';
|
||||
import { CopilotProviderType, ModelInputType, ModelOutputType } from '../types';
|
||||
import { GeminiProvider } from './gemini';
|
||||
|
||||
@@ -29,12 +27,15 @@ export class GeminiGenerativeProvider extends GeminiProvider<GeminiGenerativeCon
|
||||
ModelInputType.Text,
|
||||
ModelInputType.Image,
|
||||
ModelInputType.Audio,
|
||||
ModelInputType.File,
|
||||
],
|
||||
output: [
|
||||
ModelOutputType.Text,
|
||||
ModelOutputType.Object,
|
||||
ModelOutputType.Structured,
|
||||
],
|
||||
attachments: GEMINI_ATTACHMENT_CAPABILITY,
|
||||
structuredAttachments: GEMINI_ATTACHMENT_CAPABILITY,
|
||||
},
|
||||
],
|
||||
},
|
||||
@@ -47,12 +48,15 @@ export class GeminiGenerativeProvider extends GeminiProvider<GeminiGenerativeCon
|
||||
ModelInputType.Text,
|
||||
ModelInputType.Image,
|
||||
ModelInputType.Audio,
|
||||
ModelInputType.File,
|
||||
],
|
||||
output: [
|
||||
ModelOutputType.Text,
|
||||
ModelOutputType.Object,
|
||||
ModelOutputType.Structured,
|
||||
],
|
||||
attachments: GEMINI_ATTACHMENT_CAPABILITY,
|
||||
structuredAttachments: GEMINI_ATTACHMENT_CAPABILITY,
|
||||
},
|
||||
],
|
||||
},
|
||||
@@ -65,12 +69,15 @@ export class GeminiGenerativeProvider extends GeminiProvider<GeminiGenerativeCon
|
||||
ModelInputType.Text,
|
||||
ModelInputType.Image,
|
||||
ModelInputType.Audio,
|
||||
ModelInputType.File,
|
||||
],
|
||||
output: [
|
||||
ModelOutputType.Text,
|
||||
ModelOutputType.Object,
|
||||
ModelOutputType.Structured,
|
||||
],
|
||||
attachments: GEMINI_ATTACHMENT_CAPABILITY,
|
||||
structuredAttachments: GEMINI_ATTACHMENT_CAPABILITY,
|
||||
},
|
||||
],
|
||||
},
|
||||
@@ -86,21 +93,10 @@ export class GeminiGenerativeProvider extends GeminiProvider<GeminiGenerativeCon
|
||||
],
|
||||
},
|
||||
];
|
||||
|
||||
protected instance!: GoogleGenerativeAIProvider;
|
||||
|
||||
override configured(): boolean {
|
||||
return !!this.config.apiKey;
|
||||
}
|
||||
|
||||
protected override setup() {
|
||||
super.setup();
|
||||
this.instance = createGoogleGenerativeAI({
|
||||
apiKey: this.config.apiKey,
|
||||
baseURL: this.config.baseURL,
|
||||
});
|
||||
}
|
||||
|
||||
override async refreshOnlineModels() {
|
||||
try {
|
||||
const baseUrl =
|
||||
@@ -120,4 +116,15 @@ export class GeminiGenerativeProvider extends GeminiProvider<GeminiGenerativeCon
|
||||
this.logger.error('Failed to fetch available models', e);
|
||||
}
|
||||
}
|
||||
|
||||
protected override async createNativeConfig(): Promise<NativeLlmBackendConfig> {
|
||||
return {
|
||||
base_url: (
|
||||
this.config.baseURL ||
|
||||
'https://generativelanguage.googleapis.com/v1beta'
|
||||
).replace(/\/$/, ''),
|
||||
auth_token: this.config.apiKey,
|
||||
request_layer: 'gemini_api',
|
||||
};
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,14 +1,14 @@
|
||||
import {
|
||||
createVertex,
|
||||
type GoogleVertexProvider,
|
||||
type GoogleVertexProviderSettings,
|
||||
} from '@ai-sdk/google-vertex';
|
||||
|
||||
import type { NativeLlmBackendConfig } from '../../../../native';
|
||||
import { GEMINI_ATTACHMENT_CAPABILITY } from '../attachments';
|
||||
import { CopilotProviderType, ModelInputType, ModelOutputType } from '../types';
|
||||
import { getGoogleAuth, VertexModelListSchema } from '../utils';
|
||||
import {
|
||||
getGoogleAuth,
|
||||
VertexModelListSchema,
|
||||
type VertexProviderConfig,
|
||||
} from '../utils';
|
||||
import { GeminiProvider } from './gemini';
|
||||
|
||||
export type GeminiVertexConfig = GoogleVertexProviderSettings;
|
||||
export type GeminiVertexConfig = VertexProviderConfig;
|
||||
|
||||
export class GeminiVertexProvider extends GeminiProvider<GeminiVertexConfig> {
|
||||
override readonly type = CopilotProviderType.GeminiVertex;
|
||||
@@ -23,12 +23,15 @@ export class GeminiVertexProvider extends GeminiProvider<GeminiVertexConfig> {
|
||||
ModelInputType.Text,
|
||||
ModelInputType.Image,
|
||||
ModelInputType.Audio,
|
||||
ModelInputType.File,
|
||||
],
|
||||
output: [
|
||||
ModelOutputType.Text,
|
||||
ModelOutputType.Object,
|
||||
ModelOutputType.Structured,
|
||||
],
|
||||
attachments: GEMINI_ATTACHMENT_CAPABILITY,
|
||||
structuredAttachments: GEMINI_ATTACHMENT_CAPABILITY,
|
||||
},
|
||||
],
|
||||
},
|
||||
@@ -41,12 +44,15 @@ export class GeminiVertexProvider extends GeminiProvider<GeminiVertexConfig> {
|
||||
ModelInputType.Text,
|
||||
ModelInputType.Image,
|
||||
ModelInputType.Audio,
|
||||
ModelInputType.File,
|
||||
],
|
||||
output: [
|
||||
ModelOutputType.Text,
|
||||
ModelOutputType.Object,
|
||||
ModelOutputType.Structured,
|
||||
],
|
||||
attachments: GEMINI_ATTACHMENT_CAPABILITY,
|
||||
structuredAttachments: GEMINI_ATTACHMENT_CAPABILITY,
|
||||
},
|
||||
],
|
||||
},
|
||||
@@ -59,12 +65,15 @@ export class GeminiVertexProvider extends GeminiProvider<GeminiVertexConfig> {
|
||||
ModelInputType.Text,
|
||||
ModelInputType.Image,
|
||||
ModelInputType.Audio,
|
||||
ModelInputType.File,
|
||||
],
|
||||
output: [
|
||||
ModelOutputType.Text,
|
||||
ModelOutputType.Object,
|
||||
ModelOutputType.Structured,
|
||||
],
|
||||
attachments: GEMINI_ATTACHMENT_CAPABILITY,
|
||||
structuredAttachments: GEMINI_ATTACHMENT_CAPABILITY,
|
||||
},
|
||||
],
|
||||
},
|
||||
@@ -80,21 +89,13 @@ export class GeminiVertexProvider extends GeminiProvider<GeminiVertexConfig> {
|
||||
],
|
||||
},
|
||||
];
|
||||
|
||||
protected instance!: GoogleVertexProvider;
|
||||
|
||||
override configured(): boolean {
|
||||
return !!this.config.location && !!this.config.googleAuthOptions;
|
||||
}
|
||||
|
||||
protected override setup() {
|
||||
super.setup();
|
||||
this.instance = createVertex(this.config);
|
||||
}
|
||||
|
||||
override async refreshOnlineModels() {
|
||||
try {
|
||||
const { baseUrl, headers } = await getGoogleAuth(this.config, 'google');
|
||||
const { baseUrl, headers } = await this.resolveVertexAuth();
|
||||
if (baseUrl && !this.onlineModelList.length) {
|
||||
const { publisherModels } = await fetch(`${baseUrl}/models`, {
|
||||
headers: headers(),
|
||||
@@ -109,4 +110,19 @@ export class GeminiVertexProvider extends GeminiProvider<GeminiVertexConfig> {
|
||||
this.logger.error('Failed to fetch available models', e);
|
||||
}
|
||||
}
|
||||
|
||||
protected async resolveVertexAuth() {
|
||||
return await getGoogleAuth(this.config, 'google');
|
||||
}
|
||||
|
||||
protected override async createNativeConfig(): Promise<NativeLlmBackendConfig> {
|
||||
const auth = await this.resolveVertexAuth();
|
||||
const { Authorization: authHeader } = auth.headers();
|
||||
|
||||
return {
|
||||
base_url: auth.baseUrl || '',
|
||||
auth_token: authHeader.replace(/^Bearer\s+/i, ''),
|
||||
request_layer: 'gemini_vertex',
|
||||
};
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
import type { ToolSet } from 'ai';
|
||||
import { z } from 'zod';
|
||||
|
||||
import type {
|
||||
@@ -6,6 +5,11 @@ import type {
|
||||
NativeLlmStreamEvent,
|
||||
NativeLlmToolDefinition,
|
||||
} from '../../../native';
|
||||
import type {
|
||||
CopilotTool,
|
||||
CopilotToolExecuteOptions,
|
||||
CopilotToolSet,
|
||||
} from '../tools';
|
||||
|
||||
export type NativeDispatchFn = (
|
||||
request: NativeLlmRequest,
|
||||
@@ -16,6 +20,8 @@ export type NativeToolCall = {
|
||||
id: string;
|
||||
name: string;
|
||||
args: Record<string, unknown>;
|
||||
rawArgumentsText?: string;
|
||||
argumentParseError?: string;
|
||||
thought?: string;
|
||||
};
|
||||
|
||||
@@ -28,10 +34,18 @@ type ToolExecutionResult = {
|
||||
callId: string;
|
||||
name: string;
|
||||
args: Record<string, unknown>;
|
||||
rawArgumentsText?: string;
|
||||
argumentParseError?: string;
|
||||
output: unknown;
|
||||
isError?: boolean;
|
||||
};
|
||||
|
||||
type ParsedToolArguments = {
|
||||
args: Record<string, unknown>;
|
||||
rawArgumentsText?: string;
|
||||
argumentParseError?: string;
|
||||
};
|
||||
|
||||
export class ToolCallAccumulator {
|
||||
readonly #states = new Map<string, ToolCallState>();
|
||||
|
||||
@@ -51,12 +65,20 @@ export class ToolCallAccumulator {
|
||||
complete(event: Extract<NativeLlmStreamEvent, { type: 'tool_call' }>) {
|
||||
const state = this.#states.get(event.call_id);
|
||||
this.#states.delete(event.call_id);
|
||||
const parsed =
|
||||
event.arguments_text !== undefined || event.arguments_error !== undefined
|
||||
? {
|
||||
args: event.arguments ?? {},
|
||||
rawArgumentsText: event.arguments_text ?? state?.argumentsText,
|
||||
argumentParseError: event.arguments_error,
|
||||
}
|
||||
: event.arguments
|
||||
? this.parseArgs(event.arguments, state?.argumentsText)
|
||||
: this.parseJson(state?.argumentsText ?? '{}');
|
||||
return {
|
||||
id: event.call_id,
|
||||
name: event.name || state?.name || '',
|
||||
args: this.parseArgs(
|
||||
event.arguments ?? this.parseJson(state?.argumentsText ?? '{}')
|
||||
),
|
||||
...parsed,
|
||||
thought: event.thought,
|
||||
} satisfies NativeToolCall;
|
||||
}
|
||||
@@ -70,51 +92,61 @@ export class ToolCallAccumulator {
|
||||
pending.push({
|
||||
id: callId,
|
||||
name: state.name,
|
||||
args: this.parseArgs(this.parseJson(state.argumentsText)),
|
||||
...this.parseJson(state.argumentsText),
|
||||
});
|
||||
}
|
||||
this.#states.clear();
|
||||
return pending;
|
||||
}
|
||||
|
||||
private parseJson(jsonText: string): unknown {
|
||||
private parseJson(jsonText: string): ParsedToolArguments {
|
||||
if (!jsonText.trim()) {
|
||||
return {};
|
||||
return { args: {} };
|
||||
}
|
||||
try {
|
||||
return JSON.parse(jsonText);
|
||||
} catch {
|
||||
return {};
|
||||
return this.parseArgs(JSON.parse(jsonText), jsonText);
|
||||
} catch (error) {
|
||||
return {
|
||||
args: {},
|
||||
rawArgumentsText: jsonText,
|
||||
argumentParseError:
|
||||
error instanceof Error
|
||||
? error.message
|
||||
: 'Invalid tool arguments JSON',
|
||||
};
|
||||
}
|
||||
}
|
||||
|
||||
private parseArgs(value: unknown): Record<string, unknown> {
|
||||
private parseArgs(
|
||||
value: unknown,
|
||||
rawArgumentsText?: string
|
||||
): ParsedToolArguments {
|
||||
if (value && typeof value === 'object' && !Array.isArray(value)) {
|
||||
return value as Record<string, unknown>;
|
||||
return {
|
||||
args: value as Record<string, unknown>,
|
||||
rawArgumentsText,
|
||||
};
|
||||
}
|
||||
return {};
|
||||
return {
|
||||
args: {},
|
||||
rawArgumentsText,
|
||||
argumentParseError: 'Tool arguments must be a JSON object',
|
||||
};
|
||||
}
|
||||
}
|
||||
|
||||
export class ToolSchemaExtractor {
|
||||
static extract(toolSet: ToolSet): NativeLlmToolDefinition[] {
|
||||
static extract(toolSet: CopilotToolSet): NativeLlmToolDefinition[] {
|
||||
return Object.entries(toolSet).map(([name, tool]) => {
|
||||
const unknownTool = tool as Record<string, unknown>;
|
||||
const inputSchema =
|
||||
unknownTool.inputSchema ?? unknownTool.parameters ?? z.object({});
|
||||
|
||||
return {
|
||||
name,
|
||||
description:
|
||||
typeof unknownTool.description === 'string'
|
||||
? unknownTool.description
|
||||
: undefined,
|
||||
parameters: this.toJsonSchema(inputSchema),
|
||||
description: tool.description,
|
||||
parameters: this.toJsonSchema(tool.inputSchema ?? z.object({})),
|
||||
};
|
||||
});
|
||||
}
|
||||
|
||||
private static toJsonSchema(schema: unknown): Record<string, unknown> {
|
||||
static toJsonSchema(schema: unknown): Record<string, unknown> {
|
||||
if (!(schema instanceof z.ZodType)) {
|
||||
if (schema && typeof schema === 'object' && !Array.isArray(schema)) {
|
||||
return schema as Record<string, unknown>;
|
||||
@@ -228,14 +260,45 @@ export class ToolSchemaExtractor {
|
||||
export class ToolCallLoop {
|
||||
constructor(
|
||||
private readonly dispatch: NativeDispatchFn,
|
||||
private readonly tools: ToolSet,
|
||||
private readonly tools: CopilotToolSet,
|
||||
private readonly maxSteps = 20
|
||||
) {}
|
||||
|
||||
private normalizeToolExecuteOptions(
|
||||
signalOrOptions?: AbortSignal | CopilotToolExecuteOptions,
|
||||
maybeMessages?: CopilotToolExecuteOptions['messages']
|
||||
): CopilotToolExecuteOptions {
|
||||
if (
|
||||
signalOrOptions &&
|
||||
typeof signalOrOptions === 'object' &&
|
||||
'aborted' in signalOrOptions
|
||||
) {
|
||||
return {
|
||||
signal: signalOrOptions,
|
||||
messages: maybeMessages,
|
||||
};
|
||||
}
|
||||
|
||||
if (!signalOrOptions) {
|
||||
return maybeMessages ? { messages: maybeMessages } : {};
|
||||
}
|
||||
|
||||
return {
|
||||
...signalOrOptions,
|
||||
signal: signalOrOptions.signal,
|
||||
messages: signalOrOptions.messages ?? maybeMessages,
|
||||
};
|
||||
}
|
||||
|
||||
async *run(
|
||||
request: NativeLlmRequest,
|
||||
signal?: AbortSignal
|
||||
signalOrOptions?: AbortSignal | CopilotToolExecuteOptions,
|
||||
maybeMessages?: CopilotToolExecuteOptions['messages']
|
||||
): AsyncIterableIterator<NativeLlmStreamEvent> {
|
||||
const toolExecuteOptions = this.normalizeToolExecuteOptions(
|
||||
signalOrOptions,
|
||||
maybeMessages
|
||||
);
|
||||
const messages = request.messages.map(message => ({
|
||||
...message,
|
||||
content: [...message.content],
|
||||
@@ -253,7 +316,7 @@ export class ToolCallLoop {
|
||||
stream: true,
|
||||
messages,
|
||||
},
|
||||
signal
|
||||
toolExecuteOptions.signal
|
||||
)) {
|
||||
switch (event.type) {
|
||||
case 'tool_call_delta': {
|
||||
@@ -291,7 +354,10 @@ export class ToolCallLoop {
|
||||
throw new Error('ToolCallLoop max steps reached');
|
||||
}
|
||||
|
||||
const toolResults = await this.executeTools(toolCalls);
|
||||
const toolResults = await this.executeTools(
|
||||
toolCalls,
|
||||
toolExecuteOptions
|
||||
);
|
||||
|
||||
messages.push({
|
||||
role: 'assistant',
|
||||
@@ -300,6 +366,8 @@ export class ToolCallLoop {
|
||||
call_id: call.id,
|
||||
name: call.name,
|
||||
arguments: call.args,
|
||||
arguments_text: call.rawArgumentsText,
|
||||
arguments_error: call.argumentParseError,
|
||||
thought: call.thought,
|
||||
})),
|
||||
});
|
||||
@@ -311,6 +379,10 @@ export class ToolCallLoop {
|
||||
{
|
||||
type: 'tool_result',
|
||||
call_id: result.callId,
|
||||
name: result.name,
|
||||
arguments: result.args,
|
||||
arguments_text: result.rawArgumentsText,
|
||||
arguments_error: result.argumentParseError,
|
||||
output: result.output,
|
||||
is_error: result.isError,
|
||||
},
|
||||
@@ -321,6 +393,8 @@ export class ToolCallLoop {
|
||||
call_id: result.callId,
|
||||
name: result.name,
|
||||
arguments: result.args,
|
||||
arguments_text: result.rawArgumentsText,
|
||||
arguments_error: result.argumentParseError,
|
||||
output: result.output,
|
||||
is_error: result.isError,
|
||||
};
|
||||
@@ -328,24 +402,28 @@ export class ToolCallLoop {
|
||||
}
|
||||
}
|
||||
|
||||
private async executeTools(calls: NativeToolCall[]) {
|
||||
return await Promise.all(calls.map(call => this.executeTool(call)));
|
||||
private async executeTools(
|
||||
calls: NativeToolCall[],
|
||||
options: CopilotToolExecuteOptions
|
||||
) {
|
||||
return await Promise.all(
|
||||
calls.map(call => this.executeTool(call, options))
|
||||
);
|
||||
}
|
||||
|
||||
private async executeTool(
|
||||
call: NativeToolCall
|
||||
call: NativeToolCall,
|
||||
options: CopilotToolExecuteOptions
|
||||
): Promise<ToolExecutionResult> {
|
||||
const tool = this.tools[call.name] as
|
||||
| {
|
||||
execute?: (args: Record<string, unknown>) => Promise<unknown>;
|
||||
}
|
||||
| undefined;
|
||||
const tool = this.tools[call.name] as CopilotTool | undefined;
|
||||
|
||||
if (!tool?.execute) {
|
||||
return {
|
||||
callId: call.id,
|
||||
name: call.name,
|
||||
args: call.args,
|
||||
rawArgumentsText: call.rawArgumentsText,
|
||||
argumentParseError: call.argumentParseError,
|
||||
isError: true,
|
||||
output: {
|
||||
message: `Tool not found: ${call.name}`,
|
||||
@@ -353,12 +431,30 @@ export class ToolCallLoop {
|
||||
};
|
||||
}
|
||||
|
||||
try {
|
||||
const output = await tool.execute(call.args);
|
||||
if (call.argumentParseError) {
|
||||
return {
|
||||
callId: call.id,
|
||||
name: call.name,
|
||||
args: call.args,
|
||||
rawArgumentsText: call.rawArgumentsText,
|
||||
argumentParseError: call.argumentParseError,
|
||||
isError: true,
|
||||
output: {
|
||||
message: 'Invalid tool arguments JSON',
|
||||
rawArguments: call.rawArgumentsText,
|
||||
error: call.argumentParseError,
|
||||
},
|
||||
};
|
||||
}
|
||||
|
||||
try {
|
||||
const output = await tool.execute(call.args, options);
|
||||
return {
|
||||
callId: call.id,
|
||||
name: call.name,
|
||||
args: call.args,
|
||||
rawArgumentsText: call.rawArgumentsText,
|
||||
argumentParseError: call.argumentParseError,
|
||||
output: output ?? null,
|
||||
};
|
||||
} catch (error) {
|
||||
@@ -371,6 +467,8 @@ export class ToolCallLoop {
|
||||
callId: call.id,
|
||||
name: call.name,
|
||||
args: call.args,
|
||||
rawArgumentsText: call.rawArgumentsText,
|
||||
argumentParseError: call.argumentParseError,
|
||||
isError: true,
|
||||
output: {
|
||||
message: 'Tool execution failed',
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
import type { ToolSet } from 'ai';
|
||||
|
||||
import {
|
||||
CopilotProviderSideError,
|
||||
metrics,
|
||||
@@ -11,6 +9,7 @@ import {
|
||||
type NativeLlmRequest,
|
||||
} from '../../../native';
|
||||
import type { NodeTextMiddleware } from '../config';
|
||||
import type { CopilotToolSet } from '../tools';
|
||||
import { buildNativeRequest, NativeProviderAdapter } from './native';
|
||||
import { CopilotProvider } from './provider';
|
||||
import type {
|
||||
@@ -86,7 +85,7 @@ export class MorphProvider extends CopilotProvider<MorphConfig> {
|
||||
}
|
||||
|
||||
private createNativeAdapter(
|
||||
tools: ToolSet,
|
||||
tools: CopilotToolSet,
|
||||
nodeTextMiddleware?: NodeTextMiddleware[]
|
||||
) {
|
||||
return new NativeProviderAdapter(
|
||||
@@ -108,12 +107,14 @@ export class MorphProvider extends CopilotProvider<MorphConfig> {
|
||||
messages: PromptMessage[],
|
||||
options: CopilotChatOptions = {}
|
||||
): Promise<string> {
|
||||
const fullCond = {
|
||||
...cond,
|
||||
outputType: ModelOutputType.Text,
|
||||
};
|
||||
await this.checkParams({ messages, cond: fullCond, options });
|
||||
const model = this.selectModel(fullCond);
|
||||
const fullCond = { ...cond, outputType: ModelOutputType.Text };
|
||||
const model = this.selectModel(
|
||||
await this.checkParams({
|
||||
messages,
|
||||
cond: fullCond,
|
||||
options,
|
||||
})
|
||||
);
|
||||
|
||||
try {
|
||||
metrics.ai.counter('chat_text_calls').add(1, this.metricLabels(model.id));
|
||||
@@ -127,7 +128,7 @@ export class MorphProvider extends CopilotProvider<MorphConfig> {
|
||||
middleware,
|
||||
});
|
||||
const adapter = this.createNativeAdapter(tools, middleware.node?.text);
|
||||
return await adapter.text(request, options.signal);
|
||||
return await adapter.text(request, options.signal, messages);
|
||||
} catch (e: any) {
|
||||
metrics.ai
|
||||
.counter('chat_text_errors')
|
||||
@@ -141,12 +142,14 @@ export class MorphProvider extends CopilotProvider<MorphConfig> {
|
||||
messages: PromptMessage[],
|
||||
options: CopilotChatOptions = {}
|
||||
): AsyncIterable<string> {
|
||||
const fullCond = {
|
||||
...cond,
|
||||
outputType: ModelOutputType.Text,
|
||||
};
|
||||
await this.checkParams({ messages, cond: fullCond, options });
|
||||
const model = this.selectModel(fullCond);
|
||||
const fullCond = { ...cond, outputType: ModelOutputType.Text };
|
||||
const model = this.selectModel(
|
||||
await this.checkParams({
|
||||
messages,
|
||||
cond: fullCond,
|
||||
options,
|
||||
})
|
||||
);
|
||||
|
||||
try {
|
||||
metrics.ai
|
||||
@@ -162,7 +165,11 @@ export class MorphProvider extends CopilotProvider<MorphConfig> {
|
||||
middleware,
|
||||
});
|
||||
const adapter = this.createNativeAdapter(tools, middleware.node?.text);
|
||||
for await (const chunk of adapter.streamText(request, options.signal)) {
|
||||
for await (const chunk of adapter.streamText(
|
||||
request,
|
||||
options.signal,
|
||||
messages
|
||||
)) {
|
||||
yield chunk;
|
||||
}
|
||||
} catch (e: any) {
|
||||
|
||||
@@ -1,31 +1,41 @@
|
||||
import type { ToolSet } from 'ai';
|
||||
import { ZodType } from 'zod';
|
||||
|
||||
import { CopilotPromptInvalid } from '../../../base';
|
||||
import type {
|
||||
NativeLlmCoreContent,
|
||||
NativeLlmCoreMessage,
|
||||
NativeLlmEmbeddingRequest,
|
||||
NativeLlmRequest,
|
||||
NativeLlmStreamEvent,
|
||||
NativeLlmStructuredRequest,
|
||||
NativeLlmStructuredResponse,
|
||||
} from '../../../native';
|
||||
import type { NodeTextMiddleware, ProviderMiddlewareConfig } from '../config';
|
||||
import { NativeDispatchFn, ToolCallLoop, ToolSchemaExtractor } from './loop';
|
||||
import type { CopilotChatOptions, PromptMessage, StreamObject } from './types';
|
||||
import type { CopilotToolSet } from '../tools';
|
||||
import {
|
||||
CitationFootnoteFormatter,
|
||||
inferMimeType,
|
||||
TextStreamParser,
|
||||
} from './utils';
|
||||
|
||||
const SIMPLE_IMAGE_URL_REGEX = /^(https?:\/\/|data:image\/)/;
|
||||
canonicalizePromptAttachment,
|
||||
type CanonicalPromptAttachment,
|
||||
} from './attachments';
|
||||
import { NativeDispatchFn, ToolCallLoop, ToolSchemaExtractor } from './loop';
|
||||
import type {
|
||||
CopilotChatOptions,
|
||||
CopilotStructuredOptions,
|
||||
ModelAttachmentCapability,
|
||||
PromptMessage,
|
||||
StreamObject,
|
||||
} from './types';
|
||||
import { CitationFootnoteFormatter, TextStreamParser } from './utils';
|
||||
|
||||
type BuildNativeRequestOptions = {
|
||||
model: string;
|
||||
messages: PromptMessage[];
|
||||
options?: CopilotChatOptions;
|
||||
tools?: ToolSet;
|
||||
options?: CopilotChatOptions | CopilotStructuredOptions;
|
||||
tools?: CopilotToolSet;
|
||||
withAttachment?: boolean;
|
||||
attachmentCapability?: ModelAttachmentCapability;
|
||||
include?: string[];
|
||||
reasoning?: Record<string, unknown>;
|
||||
responseSchema?: unknown;
|
||||
middleware?: ProviderMiddlewareConfig;
|
||||
};
|
||||
|
||||
@@ -34,6 +44,11 @@ type BuildNativeRequestResult = {
|
||||
schema?: ZodType;
|
||||
};
|
||||
|
||||
type BuildNativeStructuredRequestResult = {
|
||||
request: NativeLlmStructuredRequest;
|
||||
schema: ZodType;
|
||||
};
|
||||
|
||||
type ToolCallMeta = {
|
||||
name: string;
|
||||
args: Record<string, unknown>;
|
||||
@@ -68,9 +83,121 @@ function roleToCore(role: PromptMessage['role']) {
|
||||
}
|
||||
}
|
||||
|
||||
function ensureAttachmentSupported(
|
||||
attachment: CanonicalPromptAttachment,
|
||||
attachmentCapability?: ModelAttachmentCapability
|
||||
) {
|
||||
if (!attachmentCapability) return;
|
||||
|
||||
if (!attachmentCapability.kinds.includes(attachment.kind)) {
|
||||
throw new CopilotPromptInvalid(
|
||||
`Native path does not support ${attachment.kind} attachments${
|
||||
attachment.mediaType ? ` (${attachment.mediaType})` : ''
|
||||
}`
|
||||
);
|
||||
}
|
||||
|
||||
if (
|
||||
attachmentCapability.sourceKinds?.length &&
|
||||
!attachmentCapability.sourceKinds.includes(attachment.sourceKind)
|
||||
) {
|
||||
throw new CopilotPromptInvalid(
|
||||
`Native path does not support ${attachment.sourceKind} attachment sources`
|
||||
);
|
||||
}
|
||||
|
||||
if (attachment.isRemote && attachmentCapability.allowRemoteUrls === false) {
|
||||
throw new CopilotPromptInvalid(
|
||||
'Native path does not support remote attachment urls'
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
function resolveResponseSchema(
|
||||
systemMessage: PromptMessage | undefined,
|
||||
responseSchema?: unknown
|
||||
): ZodType | undefined {
|
||||
if (responseSchema instanceof ZodType) {
|
||||
return responseSchema;
|
||||
}
|
||||
|
||||
if (systemMessage?.responseFormat?.schema instanceof ZodType) {
|
||||
return systemMessage.responseFormat.schema;
|
||||
}
|
||||
|
||||
return systemMessage?.params?.schema instanceof ZodType
|
||||
? systemMessage.params.schema
|
||||
: undefined;
|
||||
}
|
||||
|
||||
function resolveResponseStrict(
|
||||
systemMessage: PromptMessage | undefined,
|
||||
options?: CopilotStructuredOptions
|
||||
) {
|
||||
return options?.strict ?? systemMessage?.responseFormat?.strict ?? true;
|
||||
}
|
||||
|
||||
export class StructuredResponseParseError extends Error {}
|
||||
|
||||
function normalizeStructuredText(text: string) {
|
||||
const trimmed = text.replaceAll(/^ny\n/g, ' ').trim();
|
||||
if (trimmed.startsWith('```') || trimmed.endsWith('```')) {
|
||||
return trimmed
|
||||
.replace(/```[\w\s-]*\n/g, '')
|
||||
.replace(/\n```/g, '')
|
||||
.trim();
|
||||
}
|
||||
return trimmed;
|
||||
}
|
||||
|
||||
export function parseNativeStructuredOutput(
|
||||
response: Pick<NativeLlmStructuredResponse, 'output_text'> & {
|
||||
output_json?: unknown;
|
||||
}
|
||||
) {
|
||||
if (response.output_json !== undefined) {
|
||||
return response.output_json;
|
||||
}
|
||||
|
||||
const normalized = normalizeStructuredText(response.output_text);
|
||||
const candidates = [
|
||||
() => normalized,
|
||||
() => {
|
||||
const objectStart = normalized.indexOf('{');
|
||||
const objectEnd = normalized.lastIndexOf('}');
|
||||
return objectStart !== -1 && objectEnd > objectStart
|
||||
? normalized.slice(objectStart, objectEnd + 1)
|
||||
: null;
|
||||
},
|
||||
() => {
|
||||
const arrayStart = normalized.indexOf('[');
|
||||
const arrayEnd = normalized.lastIndexOf(']');
|
||||
return arrayStart !== -1 && arrayEnd > arrayStart
|
||||
? normalized.slice(arrayStart, arrayEnd + 1)
|
||||
: null;
|
||||
},
|
||||
];
|
||||
|
||||
for (const candidate of candidates) {
|
||||
try {
|
||||
const candidateText = candidate();
|
||||
if (typeof candidateText === 'string') {
|
||||
return JSON.parse(candidateText);
|
||||
}
|
||||
} catch {
|
||||
continue;
|
||||
}
|
||||
}
|
||||
|
||||
throw new StructuredResponseParseError(
|
||||
`Unexpected structured response: ${normalized.slice(0, 200)}`
|
||||
);
|
||||
}
|
||||
|
||||
async function toCoreContents(
|
||||
message: PromptMessage,
|
||||
withAttachment: boolean
|
||||
withAttachment: boolean,
|
||||
attachmentCapability?: ModelAttachmentCapability
|
||||
): Promise<NativeLlmCoreContent[]> {
|
||||
const contents: NativeLlmCoreContent[] = [];
|
||||
|
||||
@@ -81,24 +208,12 @@ async function toCoreContents(
|
||||
if (!withAttachment || !Array.isArray(message.attachments)) return contents;
|
||||
|
||||
for (const entry of message.attachments) {
|
||||
let attachmentUrl: string;
|
||||
let mediaType: string;
|
||||
|
||||
if (typeof entry === 'string') {
|
||||
attachmentUrl = entry;
|
||||
mediaType =
|
||||
typeof message.params?.mimetype === 'string'
|
||||
? message.params.mimetype
|
||||
: await inferMimeType(entry);
|
||||
} else {
|
||||
attachmentUrl = entry.attachment;
|
||||
mediaType = entry.mimeType;
|
||||
}
|
||||
|
||||
if (!SIMPLE_IMAGE_URL_REGEX.test(attachmentUrl)) continue;
|
||||
if (!mediaType.startsWith('image/')) continue;
|
||||
|
||||
contents.push({ type: 'image', source: { url: attachmentUrl } });
|
||||
const normalized = await canonicalizePromptAttachment(entry, message);
|
||||
ensureAttachmentSupported(normalized, attachmentCapability);
|
||||
contents.push({
|
||||
type: normalized.kind,
|
||||
source: normalized.source,
|
||||
});
|
||||
}
|
||||
|
||||
return contents;
|
||||
@@ -110,8 +225,10 @@ export async function buildNativeRequest({
|
||||
options = {},
|
||||
tools = {},
|
||||
withAttachment = true,
|
||||
attachmentCapability,
|
||||
include,
|
||||
reasoning,
|
||||
responseSchema,
|
||||
middleware,
|
||||
}: BuildNativeRequestOptions): Promise<BuildNativeRequestResult> {
|
||||
const copiedMessages = messages.map(message => ({
|
||||
@@ -123,10 +240,7 @@ export async function buildNativeRequest({
|
||||
|
||||
const systemMessage =
|
||||
copiedMessages[0]?.role === 'system' ? copiedMessages.shift() : undefined;
|
||||
const schema =
|
||||
systemMessage?.params?.schema instanceof ZodType
|
||||
? systemMessage.params.schema
|
||||
: undefined;
|
||||
const schema = resolveResponseSchema(systemMessage, responseSchema);
|
||||
|
||||
const coreMessages: NativeLlmCoreMessage[] = [];
|
||||
if (systemMessage?.content?.length) {
|
||||
@@ -138,7 +252,11 @@ export async function buildNativeRequest({
|
||||
|
||||
for (const message of copiedMessages) {
|
||||
if (message.role === 'system') continue;
|
||||
const content = await toCoreContents(message, withAttachment);
|
||||
const content = await toCoreContents(
|
||||
message,
|
||||
withAttachment,
|
||||
attachmentCapability
|
||||
);
|
||||
coreMessages.push({ role: roleToCore(message.role), content });
|
||||
}
|
||||
|
||||
@@ -153,6 +271,9 @@ export async function buildNativeRequest({
|
||||
tool_choice: Object.keys(tools).length ? 'auto' : undefined,
|
||||
include,
|
||||
reasoning,
|
||||
response_schema: schema
|
||||
? ToolSchemaExtractor.toJsonSchema(schema)
|
||||
: undefined,
|
||||
middleware: middleware?.rust
|
||||
? { request: middleware.rust.request, stream: middleware.rust.stream }
|
||||
: undefined,
|
||||
@@ -161,6 +282,90 @@ export async function buildNativeRequest({
|
||||
};
|
||||
}
|
||||
|
||||
export async function buildNativeStructuredRequest({
|
||||
model,
|
||||
messages,
|
||||
options = {},
|
||||
withAttachment = true,
|
||||
attachmentCapability,
|
||||
reasoning,
|
||||
responseSchema,
|
||||
middleware,
|
||||
}: Omit<
|
||||
BuildNativeRequestOptions,
|
||||
'tools' | 'include'
|
||||
>): Promise<BuildNativeStructuredRequestResult> {
|
||||
const copiedMessages = messages.map(message => ({
|
||||
...message,
|
||||
attachments: message.attachments
|
||||
? [...message.attachments]
|
||||
: message.attachments,
|
||||
}));
|
||||
|
||||
const systemMessage =
|
||||
copiedMessages[0]?.role === 'system' ? copiedMessages.shift() : undefined;
|
||||
const schema = resolveResponseSchema(systemMessage, responseSchema);
|
||||
const strict = resolveResponseStrict(systemMessage, options);
|
||||
|
||||
if (!schema) {
|
||||
throw new CopilotPromptInvalid('Schema is required');
|
||||
}
|
||||
|
||||
const coreMessages: NativeLlmCoreMessage[] = [];
|
||||
if (systemMessage?.content?.length) {
|
||||
coreMessages.push({
|
||||
role: 'system',
|
||||
content: [{ type: 'text', text: systemMessage.content }],
|
||||
});
|
||||
}
|
||||
|
||||
for (const message of copiedMessages) {
|
||||
if (message.role === 'system') continue;
|
||||
const content = await toCoreContents(
|
||||
message,
|
||||
withAttachment,
|
||||
attachmentCapability
|
||||
);
|
||||
coreMessages.push({ role: roleToCore(message.role), content });
|
||||
}
|
||||
|
||||
return {
|
||||
request: {
|
||||
model,
|
||||
messages: coreMessages,
|
||||
schema: ToolSchemaExtractor.toJsonSchema(schema),
|
||||
max_tokens: options.maxTokens ?? undefined,
|
||||
temperature: options.temperature ?? undefined,
|
||||
reasoning,
|
||||
strict,
|
||||
response_mime_type: 'application/json',
|
||||
middleware: middleware?.rust
|
||||
? { request: middleware.rust.request }
|
||||
: undefined,
|
||||
},
|
||||
schema,
|
||||
};
|
||||
}
|
||||
|
||||
export function buildNativeEmbeddingRequest({
|
||||
model,
|
||||
inputs,
|
||||
dimensions,
|
||||
taskType = 'RETRIEVAL_DOCUMENT',
|
||||
}: {
|
||||
model: string;
|
||||
inputs: string[];
|
||||
dimensions?: number;
|
||||
taskType?: string;
|
||||
}): NativeLlmEmbeddingRequest {
|
||||
return {
|
||||
model,
|
||||
inputs,
|
||||
dimensions,
|
||||
task_type: taskType,
|
||||
};
|
||||
}
|
||||
|
||||
function ensureToolResultMeta(
|
||||
event: Extract<NativeLlmStreamEvent, { type: 'tool_result' }>,
|
||||
toolCalls: Map<string, ToolCallMeta>
|
||||
@@ -244,7 +449,7 @@ export class NativeProviderAdapter {
|
||||
|
||||
constructor(
|
||||
dispatch: NativeDispatchFn,
|
||||
tools: ToolSet,
|
||||
tools: CopilotToolSet,
|
||||
maxSteps = 20,
|
||||
options: NativeProviderAdapterOptions = {}
|
||||
) {
|
||||
@@ -259,9 +464,13 @@ export class NativeProviderAdapter {
|
||||
enabledNodeTextMiddlewares.has('citation_footnote');
|
||||
}
|
||||
|
||||
async text(request: NativeLlmRequest, signal?: AbortSignal) {
|
||||
async text(
|
||||
request: NativeLlmRequest,
|
||||
signal?: AbortSignal,
|
||||
messages?: PromptMessage[]
|
||||
) {
|
||||
let output = '';
|
||||
for await (const chunk of this.streamText(request, signal)) {
|
||||
for await (const chunk of this.streamText(request, signal, messages)) {
|
||||
output += chunk;
|
||||
}
|
||||
return output.trim();
|
||||
@@ -269,7 +478,8 @@ export class NativeProviderAdapter {
|
||||
|
||||
async *streamText(
|
||||
request: NativeLlmRequest,
|
||||
signal?: AbortSignal
|
||||
signal?: AbortSignal,
|
||||
messages?: PromptMessage[]
|
||||
): AsyncIterableIterator<string> {
|
||||
const textParser = this.#enableCallout ? new TextStreamParser() : null;
|
||||
const citationFormatter = this.#enableCitationFootnote
|
||||
@@ -278,7 +488,7 @@ export class NativeProviderAdapter {
|
||||
const toolCalls = new Map<string, ToolCallMeta>();
|
||||
let streamPartId = 0;
|
||||
|
||||
for await (const event of this.#loop.run(request, signal)) {
|
||||
for await (const event of this.#loop.run(request, signal, messages)) {
|
||||
switch (event.type) {
|
||||
case 'text_delta': {
|
||||
if (textParser) {
|
||||
@@ -364,7 +574,8 @@ export class NativeProviderAdapter {
|
||||
|
||||
async *streamObject(
|
||||
request: NativeLlmRequest,
|
||||
signal?: AbortSignal
|
||||
signal?: AbortSignal,
|
||||
messages?: PromptMessage[]
|
||||
): AsyncIterableIterator<StreamObject> {
|
||||
const toolCalls = new Map<string, ToolCallMeta>();
|
||||
const citationFormatter = this.#enableCitationFootnote
|
||||
@@ -373,7 +584,7 @@ export class NativeProviderAdapter {
|
||||
const fallbackAttachmentFootnotes = new Map<string, AttachmentFootnote>();
|
||||
let hasFootnoteReference = false;
|
||||
|
||||
for await (const event of this.#loop.run(request, signal)) {
|
||||
for await (const event of this.#loop.run(request, signal, messages)) {
|
||||
switch (event.type) {
|
||||
case 'text_delta': {
|
||||
if (event.text.includes('[^')) {
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
import type { Tool, ToolSet } from 'ai';
|
||||
import { z } from 'zod';
|
||||
|
||||
import {
|
||||
@@ -12,30 +11,41 @@ import {
|
||||
} from '../../../base';
|
||||
import {
|
||||
llmDispatchStream,
|
||||
llmEmbeddingDispatch,
|
||||
llmRerankDispatch,
|
||||
llmStructuredDispatch,
|
||||
type NativeLlmBackendConfig,
|
||||
type NativeLlmEmbeddingRequest,
|
||||
type NativeLlmRequest,
|
||||
type NativeLlmRerankRequest,
|
||||
type NativeLlmRerankResponse,
|
||||
type NativeLlmStructuredRequest,
|
||||
} from '../../../native';
|
||||
import type { NodeTextMiddleware } from '../config';
|
||||
import { buildNativeRequest, NativeProviderAdapter } from './native';
|
||||
import { CopilotProvider } from './provider';
|
||||
import type { CopilotTool, CopilotToolSet } from '../tools';
|
||||
import { IMAGE_ATTACHMENT_CAPABILITY } from './attachments';
|
||||
import {
|
||||
normalizeRerankModel,
|
||||
OPENAI_RERANK_MAX_COMPLETION_TOKENS,
|
||||
OPENAI_RERANK_TOP_LOGPROBS_LIMIT,
|
||||
usesRerankReasoning,
|
||||
} from './rerank';
|
||||
buildNativeEmbeddingRequest,
|
||||
buildNativeRequest,
|
||||
buildNativeStructuredRequest,
|
||||
NativeProviderAdapter,
|
||||
parseNativeStructuredOutput,
|
||||
} from './native';
|
||||
import { CopilotProvider } from './provider';
|
||||
import type {
|
||||
CopilotChatOptions,
|
||||
CopilotChatTools,
|
||||
CopilotEmbeddingOptions,
|
||||
CopilotImageOptions,
|
||||
CopilotRerankRequest,
|
||||
CopilotStructuredOptions,
|
||||
ModelCapability,
|
||||
ModelConditions,
|
||||
PromptMessage,
|
||||
StreamObject,
|
||||
} from './types';
|
||||
import { CopilotProviderType, ModelInputType, ModelOutputType } from './types';
|
||||
import { chatToGPTMessage } from './utils';
|
||||
import { promptAttachmentToUrl } from './utils';
|
||||
|
||||
export const DEFAULT_DIMENSIONS = 256;
|
||||
|
||||
@@ -91,19 +101,6 @@ const ImageResponseSchema = z.union([
|
||||
}),
|
||||
}),
|
||||
]);
|
||||
const LogProbsSchema = z.array(
|
||||
z.object({
|
||||
token: z.string(),
|
||||
logprob: z.number(),
|
||||
top_logprobs: z.array(
|
||||
z.object({
|
||||
token: z.string(),
|
||||
logprob: z.number(),
|
||||
})
|
||||
),
|
||||
})
|
||||
);
|
||||
|
||||
const TRUSTED_ATTACHMENT_HOST_SUFFIXES = ['cdn.affine.pro'];
|
||||
|
||||
function normalizeImageFormatToMime(format?: string) {
|
||||
@@ -136,6 +133,34 @@ function normalizeImageResponseData(
|
||||
.filter((value): value is string => typeof value === 'string');
|
||||
}
|
||||
|
||||
function buildOpenAIRerankRequest(
|
||||
model: string,
|
||||
request: CopilotRerankRequest
|
||||
): NativeLlmRerankRequest {
|
||||
return {
|
||||
model,
|
||||
query: request.query,
|
||||
candidates: request.candidates.map(candidate => ({
|
||||
...(candidate.id ? { id: candidate.id } : {}),
|
||||
text: candidate.text,
|
||||
})),
|
||||
...(request.topK ? { top_n: request.topK } : {}),
|
||||
};
|
||||
}
|
||||
|
||||
function createOpenAIMultimodalCapability(
|
||||
output: ModelCapability['output'],
|
||||
options: Pick<ModelCapability, 'defaultForOutputType'> = {}
|
||||
): ModelCapability {
|
||||
return {
|
||||
input: [ModelInputType.Text, ModelInputType.Image],
|
||||
output,
|
||||
attachments: IMAGE_ATTACHMENT_CAPABILITY,
|
||||
structuredAttachments: IMAGE_ATTACHMENT_CAPABILITY,
|
||||
...options,
|
||||
};
|
||||
}
|
||||
|
||||
export class OpenAIProvider extends CopilotProvider<OpenAIConfig> {
|
||||
readonly type = CopilotProviderType.OpenAI;
|
||||
|
||||
@@ -145,10 +170,10 @@ export class OpenAIProvider extends CopilotProvider<OpenAIConfig> {
|
||||
name: 'GPT 4o',
|
||||
id: 'gpt-4o',
|
||||
capabilities: [
|
||||
{
|
||||
input: [ModelInputType.Text, ModelInputType.Image],
|
||||
output: [ModelOutputType.Text, ModelOutputType.Object],
|
||||
},
|
||||
createOpenAIMultimodalCapability([
|
||||
ModelOutputType.Text,
|
||||
ModelOutputType.Object,
|
||||
]),
|
||||
],
|
||||
},
|
||||
// FIXME(@darkskygit): deprecated
|
||||
@@ -156,20 +181,20 @@ export class OpenAIProvider extends CopilotProvider<OpenAIConfig> {
|
||||
name: 'GPT 4o 2024-08-06',
|
||||
id: 'gpt-4o-2024-08-06',
|
||||
capabilities: [
|
||||
{
|
||||
input: [ModelInputType.Text, ModelInputType.Image],
|
||||
output: [ModelOutputType.Text, ModelOutputType.Object],
|
||||
},
|
||||
createOpenAIMultimodalCapability([
|
||||
ModelOutputType.Text,
|
||||
ModelOutputType.Object,
|
||||
]),
|
||||
],
|
||||
},
|
||||
{
|
||||
name: 'GPT 4o Mini',
|
||||
id: 'gpt-4o-mini',
|
||||
capabilities: [
|
||||
{
|
||||
input: [ModelInputType.Text, ModelInputType.Image],
|
||||
output: [ModelOutputType.Text, ModelOutputType.Object],
|
||||
},
|
||||
createOpenAIMultimodalCapability([
|
||||
ModelOutputType.Text,
|
||||
ModelOutputType.Object,
|
||||
]),
|
||||
],
|
||||
},
|
||||
// FIXME(@darkskygit): deprecated
|
||||
@@ -177,181 +202,158 @@ export class OpenAIProvider extends CopilotProvider<OpenAIConfig> {
|
||||
name: 'GPT 4o Mini 2024-07-18',
|
||||
id: 'gpt-4o-mini-2024-07-18',
|
||||
capabilities: [
|
||||
{
|
||||
input: [ModelInputType.Text, ModelInputType.Image],
|
||||
output: [ModelOutputType.Text, ModelOutputType.Object],
|
||||
},
|
||||
createOpenAIMultimodalCapability([
|
||||
ModelOutputType.Text,
|
||||
ModelOutputType.Object,
|
||||
]),
|
||||
],
|
||||
},
|
||||
{
|
||||
name: 'GPT 4.1',
|
||||
id: 'gpt-4.1',
|
||||
capabilities: [
|
||||
{
|
||||
input: [ModelInputType.Text, ModelInputType.Image],
|
||||
output: [
|
||||
createOpenAIMultimodalCapability(
|
||||
[
|
||||
ModelOutputType.Text,
|
||||
ModelOutputType.Object,
|
||||
ModelOutputType.Rerank,
|
||||
ModelOutputType.Structured,
|
||||
],
|
||||
defaultForOutputType: true,
|
||||
},
|
||||
{ defaultForOutputType: true }
|
||||
),
|
||||
],
|
||||
},
|
||||
{
|
||||
name: 'GPT 4.1 2025-04-14',
|
||||
id: 'gpt-4.1-2025-04-14',
|
||||
capabilities: [
|
||||
{
|
||||
input: [ModelInputType.Text, ModelInputType.Image],
|
||||
output: [
|
||||
ModelOutputType.Text,
|
||||
ModelOutputType.Object,
|
||||
ModelOutputType.Structured,
|
||||
],
|
||||
},
|
||||
createOpenAIMultimodalCapability([
|
||||
ModelOutputType.Text,
|
||||
ModelOutputType.Object,
|
||||
ModelOutputType.Rerank,
|
||||
ModelOutputType.Structured,
|
||||
]),
|
||||
],
|
||||
},
|
||||
{
|
||||
name: 'GPT 4.1 Mini',
|
||||
id: 'gpt-4.1-mini',
|
||||
capabilities: [
|
||||
{
|
||||
input: [ModelInputType.Text, ModelInputType.Image],
|
||||
output: [
|
||||
ModelOutputType.Text,
|
||||
ModelOutputType.Object,
|
||||
ModelOutputType.Structured,
|
||||
],
|
||||
},
|
||||
createOpenAIMultimodalCapability([
|
||||
ModelOutputType.Text,
|
||||
ModelOutputType.Object,
|
||||
ModelOutputType.Rerank,
|
||||
ModelOutputType.Structured,
|
||||
]),
|
||||
],
|
||||
},
|
||||
{
|
||||
name: 'GPT 4.1 Nano',
|
||||
id: 'gpt-4.1-nano',
|
||||
capabilities: [
|
||||
{
|
||||
input: [ModelInputType.Text, ModelInputType.Image],
|
||||
output: [
|
||||
ModelOutputType.Text,
|
||||
ModelOutputType.Object,
|
||||
ModelOutputType.Structured,
|
||||
],
|
||||
},
|
||||
createOpenAIMultimodalCapability([
|
||||
ModelOutputType.Text,
|
||||
ModelOutputType.Object,
|
||||
ModelOutputType.Rerank,
|
||||
ModelOutputType.Structured,
|
||||
]),
|
||||
],
|
||||
},
|
||||
{
|
||||
name: 'GPT 5',
|
||||
id: 'gpt-5',
|
||||
capabilities: [
|
||||
{
|
||||
input: [ModelInputType.Text, ModelInputType.Image],
|
||||
output: [
|
||||
ModelOutputType.Text,
|
||||
ModelOutputType.Object,
|
||||
ModelOutputType.Structured,
|
||||
],
|
||||
},
|
||||
createOpenAIMultimodalCapability([
|
||||
ModelOutputType.Text,
|
||||
ModelOutputType.Object,
|
||||
ModelOutputType.Structured,
|
||||
]),
|
||||
],
|
||||
},
|
||||
{
|
||||
name: 'GPT 5 2025-08-07',
|
||||
id: 'gpt-5-2025-08-07',
|
||||
capabilities: [
|
||||
{
|
||||
input: [ModelInputType.Text, ModelInputType.Image],
|
||||
output: [
|
||||
ModelOutputType.Text,
|
||||
ModelOutputType.Object,
|
||||
ModelOutputType.Structured,
|
||||
],
|
||||
},
|
||||
createOpenAIMultimodalCapability([
|
||||
ModelOutputType.Text,
|
||||
ModelOutputType.Object,
|
||||
ModelOutputType.Structured,
|
||||
]),
|
||||
],
|
||||
},
|
||||
{
|
||||
name: 'GPT 5 Mini',
|
||||
id: 'gpt-5-mini',
|
||||
capabilities: [
|
||||
{
|
||||
input: [ModelInputType.Text, ModelInputType.Image],
|
||||
output: [
|
||||
ModelOutputType.Text,
|
||||
ModelOutputType.Object,
|
||||
ModelOutputType.Structured,
|
||||
],
|
||||
},
|
||||
createOpenAIMultimodalCapability([
|
||||
ModelOutputType.Text,
|
||||
ModelOutputType.Object,
|
||||
ModelOutputType.Structured,
|
||||
]),
|
||||
],
|
||||
},
|
||||
{
|
||||
name: 'GPT 5.2',
|
||||
id: 'gpt-5.2',
|
||||
capabilities: [
|
||||
{
|
||||
input: [ModelInputType.Text, ModelInputType.Image],
|
||||
output: [
|
||||
ModelOutputType.Text,
|
||||
ModelOutputType.Object,
|
||||
ModelOutputType.Structured,
|
||||
],
|
||||
},
|
||||
createOpenAIMultimodalCapability([
|
||||
ModelOutputType.Text,
|
||||
ModelOutputType.Object,
|
||||
ModelOutputType.Rerank,
|
||||
ModelOutputType.Structured,
|
||||
]),
|
||||
],
|
||||
},
|
||||
{
|
||||
name: 'GPT 5.2 2025-12-11',
|
||||
id: 'gpt-5.2-2025-12-11',
|
||||
capabilities: [
|
||||
{
|
||||
input: [ModelInputType.Text, ModelInputType.Image],
|
||||
output: [
|
||||
ModelOutputType.Text,
|
||||
ModelOutputType.Object,
|
||||
ModelOutputType.Structured,
|
||||
],
|
||||
},
|
||||
createOpenAIMultimodalCapability([
|
||||
ModelOutputType.Text,
|
||||
ModelOutputType.Object,
|
||||
ModelOutputType.Structured,
|
||||
]),
|
||||
],
|
||||
},
|
||||
{
|
||||
name: 'GPT 5 Nano',
|
||||
id: 'gpt-5-nano',
|
||||
capabilities: [
|
||||
{
|
||||
input: [ModelInputType.Text, ModelInputType.Image],
|
||||
output: [
|
||||
ModelOutputType.Text,
|
||||
ModelOutputType.Object,
|
||||
ModelOutputType.Structured,
|
||||
],
|
||||
},
|
||||
createOpenAIMultimodalCapability([
|
||||
ModelOutputType.Text,
|
||||
ModelOutputType.Object,
|
||||
ModelOutputType.Structured,
|
||||
]),
|
||||
],
|
||||
},
|
||||
{
|
||||
name: 'GPT O1',
|
||||
id: 'o1',
|
||||
capabilities: [
|
||||
{
|
||||
input: [ModelInputType.Text, ModelInputType.Image],
|
||||
output: [ModelOutputType.Text, ModelOutputType.Object],
|
||||
},
|
||||
createOpenAIMultimodalCapability([
|
||||
ModelOutputType.Text,
|
||||
ModelOutputType.Object,
|
||||
]),
|
||||
],
|
||||
},
|
||||
{
|
||||
name: 'GPT O3',
|
||||
id: 'o3',
|
||||
capabilities: [
|
||||
{
|
||||
input: [ModelInputType.Text, ModelInputType.Image],
|
||||
output: [ModelOutputType.Text, ModelOutputType.Object],
|
||||
},
|
||||
createOpenAIMultimodalCapability([
|
||||
ModelOutputType.Text,
|
||||
ModelOutputType.Object,
|
||||
]),
|
||||
],
|
||||
},
|
||||
{
|
||||
name: 'GPT O4 Mini',
|
||||
id: 'o4-mini',
|
||||
capabilities: [
|
||||
{
|
||||
input: [ModelInputType.Text, ModelInputType.Image],
|
||||
output: [ModelOutputType.Text, ModelOutputType.Object],
|
||||
},
|
||||
createOpenAIMultimodalCapability([
|
||||
ModelOutputType.Text,
|
||||
ModelOutputType.Object,
|
||||
]),
|
||||
],
|
||||
},
|
||||
// Embedding models
|
||||
@@ -387,11 +389,9 @@ export class OpenAIProvider extends CopilotProvider<OpenAIConfig> {
|
||||
{
|
||||
id: 'gpt-image-1',
|
||||
capabilities: [
|
||||
{
|
||||
input: [ModelInputType.Text, ModelInputType.Image],
|
||||
output: [ModelOutputType.Image],
|
||||
createOpenAIMultimodalCapability([ModelOutputType.Image], {
|
||||
defaultForOutputType: true,
|
||||
},
|
||||
}),
|
||||
],
|
||||
},
|
||||
];
|
||||
@@ -437,7 +437,7 @@ export class OpenAIProvider extends CopilotProvider<OpenAIConfig> {
|
||||
override getProviderSpecificTools(
|
||||
toolName: CopilotChatTools,
|
||||
_model: string
|
||||
): [string, Tool?] | undefined {
|
||||
): [string, CopilotTool?] | undefined {
|
||||
if (toolName === 'docEdit') {
|
||||
return ['doc_edit', undefined];
|
||||
}
|
||||
@@ -452,14 +452,18 @@ export class OpenAIProvider extends CopilotProvider<OpenAIConfig> {
|
||||
};
|
||||
}
|
||||
|
||||
private getNativeProtocol() {
|
||||
return this.config.oldApiStyle ? 'openai_chat' : 'openai_responses';
|
||||
}
|
||||
|
||||
private createNativeAdapter(
|
||||
tools: ToolSet,
|
||||
tools: CopilotToolSet,
|
||||
nodeTextMiddleware?: NodeTextMiddleware[]
|
||||
) {
|
||||
return new NativeProviderAdapter(
|
||||
(request: NativeLlmRequest, signal?: AbortSignal) =>
|
||||
llmDispatchStream(
|
||||
this.config.oldApiStyle ? 'openai_chat' : 'openai_responses',
|
||||
this.getNativeProtocol(),
|
||||
this.createNativeConfig(),
|
||||
request,
|
||||
signal
|
||||
@@ -470,6 +474,27 @@ export class OpenAIProvider extends CopilotProvider<OpenAIConfig> {
|
||||
);
|
||||
}
|
||||
|
||||
protected createNativeStructuredDispatch(
|
||||
backendConfig: NativeLlmBackendConfig
|
||||
) {
|
||||
return (request: NativeLlmStructuredRequest) =>
|
||||
llmStructuredDispatch(this.getNativeProtocol(), backendConfig, request);
|
||||
}
|
||||
|
||||
protected createNativeEmbeddingDispatch(
|
||||
backendConfig: NativeLlmBackendConfig
|
||||
) {
|
||||
return (request: NativeLlmEmbeddingRequest) =>
|
||||
llmEmbeddingDispatch(this.getNativeProtocol(), backendConfig, request);
|
||||
}
|
||||
|
||||
protected createNativeRerankDispatch(backendConfig: NativeLlmBackendConfig) {
|
||||
return (
|
||||
request: NativeLlmRerankRequest
|
||||
): Promise<NativeLlmRerankResponse> =>
|
||||
llmRerankDispatch('openai_chat', backendConfig, request);
|
||||
}
|
||||
|
||||
private getReasoning(
|
||||
options: NonNullable<CopilotChatOptions>,
|
||||
model: string
|
||||
@@ -486,13 +511,18 @@ export class OpenAIProvider extends CopilotProvider<OpenAIConfig> {
|
||||
options: CopilotChatOptions = {}
|
||||
): Promise<string> {
|
||||
const fullCond = { ...cond, outputType: ModelOutputType.Text };
|
||||
await this.checkParams({ messages, cond: fullCond, options });
|
||||
const model = this.selectModel(fullCond);
|
||||
const normalizedCond = await this.checkParams({
|
||||
messages,
|
||||
cond: fullCond,
|
||||
options,
|
||||
});
|
||||
const model = this.selectModel(normalizedCond);
|
||||
|
||||
try {
|
||||
metrics.ai.counter('chat_text_calls').add(1, this.metricLabels(model.id));
|
||||
const tools = await this.getTools(options, model.id);
|
||||
const middleware = this.getActiveProviderMiddleware();
|
||||
const cap = this.getAttachCapability(model, ModelOutputType.Text);
|
||||
const normalizedOptions = normalizeOpenAIOptionsForModel(
|
||||
options,
|
||||
model.id
|
||||
@@ -502,12 +532,13 @@ export class OpenAIProvider extends CopilotProvider<OpenAIConfig> {
|
||||
messages,
|
||||
options: normalizedOptions,
|
||||
tools,
|
||||
attachmentCapability: cap,
|
||||
include: options.webSearch ? ['citations'] : undefined,
|
||||
reasoning: this.getReasoning(options, model.id),
|
||||
middleware,
|
||||
});
|
||||
const adapter = this.createNativeAdapter(tools, middleware.node?.text);
|
||||
return await adapter.text(request, options.signal);
|
||||
return await adapter.text(request, options.signal, messages);
|
||||
} catch (e: any) {
|
||||
metrics.ai
|
||||
.counter('chat_text_errors')
|
||||
@@ -525,8 +556,12 @@ export class OpenAIProvider extends CopilotProvider<OpenAIConfig> {
|
||||
...cond,
|
||||
outputType: ModelOutputType.Text,
|
||||
};
|
||||
await this.checkParams({ messages, cond: fullCond, options });
|
||||
const model = this.selectModel(fullCond);
|
||||
const normalizedCond = await this.checkParams({
|
||||
messages,
|
||||
cond: fullCond,
|
||||
options,
|
||||
});
|
||||
const model = this.selectModel(normalizedCond);
|
||||
|
||||
try {
|
||||
metrics.ai
|
||||
@@ -534,6 +569,7 @@ export class OpenAIProvider extends CopilotProvider<OpenAIConfig> {
|
||||
.add(1, this.metricLabels(model.id));
|
||||
const tools = await this.getTools(options, model.id);
|
||||
const middleware = this.getActiveProviderMiddleware();
|
||||
const cap = this.getAttachCapability(model, ModelOutputType.Text);
|
||||
const normalizedOptions = normalizeOpenAIOptionsForModel(
|
||||
options,
|
||||
model.id
|
||||
@@ -543,12 +579,17 @@ export class OpenAIProvider extends CopilotProvider<OpenAIConfig> {
|
||||
messages,
|
||||
options: normalizedOptions,
|
||||
tools,
|
||||
attachmentCapability: cap,
|
||||
include: options.webSearch ? ['citations'] : undefined,
|
||||
reasoning: this.getReasoning(options, model.id),
|
||||
middleware,
|
||||
});
|
||||
const adapter = this.createNativeAdapter(tools, middleware.node?.text);
|
||||
for await (const chunk of adapter.streamText(request, options.signal)) {
|
||||
for await (const chunk of adapter.streamText(
|
||||
request,
|
||||
options.signal,
|
||||
messages
|
||||
)) {
|
||||
yield chunk;
|
||||
}
|
||||
} catch (e: any) {
|
||||
@@ -565,8 +606,12 @@ export class OpenAIProvider extends CopilotProvider<OpenAIConfig> {
|
||||
options: CopilotChatOptions = {}
|
||||
): AsyncIterable<StreamObject> {
|
||||
const fullCond = { ...cond, outputType: ModelOutputType.Object };
|
||||
await this.checkParams({ cond: fullCond, messages, options });
|
||||
const model = this.selectModel(fullCond);
|
||||
const normalizedCond = await this.checkParams({
|
||||
cond: fullCond,
|
||||
messages,
|
||||
options,
|
||||
});
|
||||
const model = this.selectModel(normalizedCond);
|
||||
|
||||
try {
|
||||
metrics.ai
|
||||
@@ -574,6 +619,7 @@ export class OpenAIProvider extends CopilotProvider<OpenAIConfig> {
|
||||
.add(1, this.metricLabels(model.id));
|
||||
const tools = await this.getTools(options, model.id);
|
||||
const middleware = this.getActiveProviderMiddleware();
|
||||
const cap = this.getAttachCapability(model, ModelOutputType.Object);
|
||||
const normalizedOptions = normalizeOpenAIOptionsForModel(
|
||||
options,
|
||||
model.id
|
||||
@@ -583,12 +629,17 @@ export class OpenAIProvider extends CopilotProvider<OpenAIConfig> {
|
||||
messages,
|
||||
options: normalizedOptions,
|
||||
tools,
|
||||
attachmentCapability: cap,
|
||||
include: options.webSearch ? ['citations'] : undefined,
|
||||
reasoning: this.getReasoning(options, model.id),
|
||||
middleware,
|
||||
});
|
||||
const adapter = this.createNativeAdapter(tools, middleware.node?.text);
|
||||
for await (const chunk of adapter.streamObject(request, options.signal)) {
|
||||
for await (const chunk of adapter.streamObject(
|
||||
request,
|
||||
options.signal,
|
||||
messages
|
||||
)) {
|
||||
yield chunk;
|
||||
}
|
||||
} catch (e: any) {
|
||||
@@ -605,31 +656,34 @@ export class OpenAIProvider extends CopilotProvider<OpenAIConfig> {
|
||||
options: CopilotStructuredOptions = {}
|
||||
): Promise<string> {
|
||||
const fullCond = { ...cond, outputType: ModelOutputType.Structured };
|
||||
await this.checkParams({ messages, cond: fullCond, options });
|
||||
const model = this.selectModel(fullCond);
|
||||
const normalizedCond = await this.checkParams({
|
||||
messages,
|
||||
cond: fullCond,
|
||||
options,
|
||||
});
|
||||
const model = this.selectModel(normalizedCond);
|
||||
|
||||
try {
|
||||
metrics.ai.counter('chat_text_calls').add(1, { model: model.id });
|
||||
const tools = await this.getTools(options, model.id);
|
||||
const backendConfig = this.createNativeConfig();
|
||||
const middleware = this.getActiveProviderMiddleware();
|
||||
const cap = this.getAttachCapability(model, ModelOutputType.Structured);
|
||||
const normalizedOptions = normalizeOpenAIOptionsForModel(
|
||||
options,
|
||||
model.id
|
||||
);
|
||||
const { request, schema } = await buildNativeRequest({
|
||||
const { request, schema } = await buildNativeStructuredRequest({
|
||||
model: model.id,
|
||||
messages,
|
||||
options: normalizedOptions,
|
||||
tools,
|
||||
attachmentCapability: cap,
|
||||
reasoning: this.getReasoning(options, model.id),
|
||||
responseSchema: options.schema,
|
||||
middleware,
|
||||
});
|
||||
if (!schema) {
|
||||
throw new CopilotPromptInvalid('Schema is required');
|
||||
}
|
||||
const adapter = this.createNativeAdapter(tools, middleware.node?.text);
|
||||
const text = await adapter.text(request, options.signal);
|
||||
const parsed = JSON.parse(text);
|
||||
const response =
|
||||
await this.createNativeStructuredDispatch(backendConfig)(request);
|
||||
const parsed = parseNativeStructuredOutput(response);
|
||||
const validated = schema.parse(parsed);
|
||||
return JSON.stringify(validated);
|
||||
} catch (e: any) {
|
||||
@@ -640,71 +694,26 @@ export class OpenAIProvider extends CopilotProvider<OpenAIConfig> {
|
||||
|
||||
override async rerank(
|
||||
cond: ModelConditions,
|
||||
chunkMessages: PromptMessage[][],
|
||||
request: CopilotRerankRequest,
|
||||
options: CopilotChatOptions = {}
|
||||
): Promise<number[]> {
|
||||
const fullCond = { ...cond, outputType: ModelOutputType.Text };
|
||||
await this.checkParams({ messages: [], cond: fullCond, options });
|
||||
const model = this.selectModel(fullCond);
|
||||
const fullCond = { ...cond, outputType: ModelOutputType.Rerank };
|
||||
const normalizedCond = await this.checkParams({
|
||||
messages: [],
|
||||
cond: fullCond,
|
||||
options,
|
||||
});
|
||||
const model = this.selectModel(normalizedCond);
|
||||
|
||||
const scores = await Promise.all(
|
||||
chunkMessages.map(async messages => {
|
||||
const [system, msgs] = await chatToGPTMessage(messages);
|
||||
const rerankModel = normalizeRerankModel(model.id);
|
||||
const response = await this.requestOpenAIJson(
|
||||
'/chat/completions',
|
||||
{
|
||||
model: rerankModel,
|
||||
messages: this.toOpenAIChatMessages(system, msgs),
|
||||
temperature: 0,
|
||||
logprobs: true,
|
||||
top_logprobs: OPENAI_RERANK_TOP_LOGPROBS_LIMIT,
|
||||
...(usesRerankReasoning(rerankModel)
|
||||
? {
|
||||
reasoning_effort: 'none' as const,
|
||||
max_completion_tokens: OPENAI_RERANK_MAX_COMPLETION_TOKENS,
|
||||
}
|
||||
: { max_tokens: OPENAI_RERANK_MAX_COMPLETION_TOKENS }),
|
||||
},
|
||||
options.signal
|
||||
);
|
||||
|
||||
const logprobs = response?.choices?.[0]?.logprobs?.content;
|
||||
if (!Array.isArray(logprobs) || logprobs.length === 0) {
|
||||
return 0;
|
||||
}
|
||||
|
||||
const parsedLogprobs = LogProbsSchema.parse(logprobs);
|
||||
const topMap = parsedLogprobs[0].top_logprobs.reduce(
|
||||
(acc, { token, logprob }) => ({ ...acc, [token]: logprob }),
|
||||
{} as Record<string, number>
|
||||
);
|
||||
|
||||
const findLogProb = (token: string): number => {
|
||||
// OpenAI often includes a leading space, so try matching '.yes', '_yes', ' yes' and 'yes'
|
||||
return [...'_:. "-\t,(=_“'.split('').map(c => c + token), token]
|
||||
.flatMap(v => [v, v.toLowerCase(), v.toUpperCase()])
|
||||
.reduce<number>(
|
||||
(best, key) =>
|
||||
(topMap[key] ?? Number.NEGATIVE_INFINITY) > best
|
||||
? topMap[key]
|
||||
: best,
|
||||
Number.NEGATIVE_INFINITY
|
||||
);
|
||||
};
|
||||
|
||||
const logYes = findLogProb('Yes');
|
||||
const logNo = findLogProb('No');
|
||||
|
||||
const pYes = Math.exp(logYes);
|
||||
const pNo = Math.exp(logNo);
|
||||
const prob = pYes + pNo === 0 ? 0 : pYes / (pYes + pNo);
|
||||
|
||||
return prob;
|
||||
})
|
||||
);
|
||||
|
||||
return scores;
|
||||
try {
|
||||
const backendConfig = this.createNativeConfig();
|
||||
const nativeRequest = buildOpenAIRerankRequest(model.id, request);
|
||||
const response =
|
||||
await this.createNativeRerankDispatch(backendConfig)(nativeRequest);
|
||||
return response.scores;
|
||||
} catch (e: any) {
|
||||
throw this.handleError(e);
|
||||
}
|
||||
}
|
||||
|
||||
// ====== text to image ======
|
||||
@@ -906,7 +915,8 @@ export class OpenAIProvider extends CopilotProvider<OpenAIConfig> {
|
||||
form.set('output_format', outputFormat);
|
||||
|
||||
for (const [idx, entry] of attachments.entries()) {
|
||||
const url = typeof entry === 'string' ? entry : entry.attachment;
|
||||
const url = promptAttachmentToUrl(entry);
|
||||
if (!url) continue;
|
||||
try {
|
||||
const attachment = await this.fetchImage(url, maxBytes, signal);
|
||||
if (!attachment) continue;
|
||||
@@ -964,8 +974,12 @@ export class OpenAIProvider extends CopilotProvider<OpenAIConfig> {
|
||||
options: CopilotImageOptions = {}
|
||||
) {
|
||||
const fullCond = { ...cond, outputType: ModelOutputType.Image };
|
||||
await this.checkParams({ messages, cond: fullCond, options });
|
||||
const model = this.selectModel(fullCond);
|
||||
const normalizedCond = await this.checkParams({
|
||||
messages,
|
||||
cond: fullCond,
|
||||
options,
|
||||
});
|
||||
const model = this.selectModel(normalizedCond);
|
||||
|
||||
metrics.ai
|
||||
.counter('generate_images_stream_calls')
|
||||
@@ -1017,65 +1031,36 @@ export class OpenAIProvider extends CopilotProvider<OpenAIConfig> {
|
||||
messages: string | string[],
|
||||
options: CopilotEmbeddingOptions = { dimensions: DEFAULT_DIMENSIONS }
|
||||
): Promise<number[][]> {
|
||||
messages = Array.isArray(messages) ? messages : [messages];
|
||||
const input = Array.isArray(messages) ? messages : [messages];
|
||||
const fullCond = { ...cond, outputType: ModelOutputType.Embedding };
|
||||
await this.checkParams({ embeddings: messages, cond: fullCond, options });
|
||||
const model = this.selectModel(fullCond);
|
||||
const normalizedCond = await this.checkParams({
|
||||
embeddings: input,
|
||||
cond: fullCond,
|
||||
options,
|
||||
});
|
||||
const model = this.selectModel(normalizedCond);
|
||||
|
||||
try {
|
||||
metrics.ai
|
||||
.counter('generate_embedding_calls')
|
||||
.add(1, { model: model.id });
|
||||
const response = await this.requestOpenAIJson('/embeddings', {
|
||||
model: model.id,
|
||||
input: messages,
|
||||
dimensions: options.dimensions || DEFAULT_DIMENSIONS,
|
||||
});
|
||||
const data = Array.isArray(response?.data) ? response.data : [];
|
||||
return data
|
||||
.map((item: any) => item?.embedding)
|
||||
.filter((embedding: unknown) => Array.isArray(embedding)) as number[][];
|
||||
.add(1, this.metricLabels(model.id));
|
||||
const backendConfig = this.createNativeConfig();
|
||||
const response = await this.createNativeEmbeddingDispatch(backendConfig)(
|
||||
buildNativeEmbeddingRequest({
|
||||
model: model.id,
|
||||
inputs: input,
|
||||
dimensions: options.dimensions || DEFAULT_DIMENSIONS,
|
||||
})
|
||||
);
|
||||
return response.embeddings;
|
||||
} catch (e: any) {
|
||||
metrics.ai
|
||||
.counter('generate_embedding_errors')
|
||||
.add(1, { model: model.id });
|
||||
.add(1, this.metricLabels(model.id));
|
||||
throw this.handleError(e);
|
||||
}
|
||||
}
|
||||
|
||||
private toOpenAIChatMessages(
|
||||
system: string | undefined,
|
||||
messages: Awaited<ReturnType<typeof chatToGPTMessage>>[1]
|
||||
) {
|
||||
const result: Array<{ role: string; content: string }> = [];
|
||||
if (system) {
|
||||
result.push({ role: 'system', content: system });
|
||||
}
|
||||
|
||||
for (const message of messages) {
|
||||
if (typeof message.content === 'string') {
|
||||
result.push({ role: message.role, content: message.content });
|
||||
continue;
|
||||
}
|
||||
|
||||
const text = message.content
|
||||
.filter(
|
||||
part =>
|
||||
part &&
|
||||
typeof part === 'object' &&
|
||||
'type' in part &&
|
||||
part.type === 'text' &&
|
||||
'text' in part
|
||||
)
|
||||
.map(part => String((part as { text: string }).text))
|
||||
.join('\n');
|
||||
|
||||
result.push({ role: message.role, content: text || '[no content]' });
|
||||
}
|
||||
|
||||
return result;
|
||||
}
|
||||
|
||||
private async requestOpenAIJson(
|
||||
path: string,
|
||||
body: Record<string, unknown>,
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
import type { ToolSet } from 'ai';
|
||||
|
||||
import { CopilotProviderSideError, metrics } from '../../../base';
|
||||
import {
|
||||
llmDispatchStream,
|
||||
@@ -7,6 +5,7 @@ import {
|
||||
type NativeLlmRequest,
|
||||
} from '../../../native';
|
||||
import type { NodeTextMiddleware } from '../config';
|
||||
import type { CopilotToolSet } from '../tools';
|
||||
import { buildNativeRequest, NativeProviderAdapter } from './native';
|
||||
import { CopilotProvider } from './provider';
|
||||
import {
|
||||
@@ -87,7 +86,7 @@ export class PerplexityProvider extends CopilotProvider<PerplexityConfig> {
|
||||
}
|
||||
|
||||
private createNativeAdapter(
|
||||
tools: ToolSet,
|
||||
tools: CopilotToolSet,
|
||||
nodeTextMiddleware?: NodeTextMiddleware[]
|
||||
) {
|
||||
return new NativeProviderAdapter(
|
||||
@@ -110,8 +109,13 @@ export class PerplexityProvider extends CopilotProvider<PerplexityConfig> {
|
||||
options: CopilotChatOptions = {}
|
||||
): Promise<string> {
|
||||
const fullCond = { ...cond, outputType: ModelOutputType.Text };
|
||||
await this.checkParams({ cond: fullCond, messages, options });
|
||||
const model = this.selectModel(fullCond);
|
||||
const normalizedCond = await this.checkParams({
|
||||
cond: fullCond,
|
||||
messages,
|
||||
options,
|
||||
withAttachment: false,
|
||||
});
|
||||
const model = this.selectModel(normalizedCond);
|
||||
|
||||
try {
|
||||
metrics.ai.counter('chat_text_calls').add(1, this.metricLabels(model.id));
|
||||
@@ -128,7 +132,7 @@ export class PerplexityProvider extends CopilotProvider<PerplexityConfig> {
|
||||
middleware,
|
||||
});
|
||||
const adapter = this.createNativeAdapter(tools, middleware.node?.text);
|
||||
return await adapter.text(request, options.signal);
|
||||
return await adapter.text(request, options.signal, messages);
|
||||
} catch (e: any) {
|
||||
metrics.ai
|
||||
.counter('chat_text_errors')
|
||||
@@ -143,8 +147,13 @@ export class PerplexityProvider extends CopilotProvider<PerplexityConfig> {
|
||||
options: CopilotChatOptions = {}
|
||||
): AsyncIterable<string> {
|
||||
const fullCond = { ...cond, outputType: ModelOutputType.Text };
|
||||
await this.checkParams({ cond: fullCond, messages, options });
|
||||
const model = this.selectModel(fullCond);
|
||||
const normalizedCond = await this.checkParams({
|
||||
cond: fullCond,
|
||||
messages,
|
||||
options,
|
||||
withAttachment: false,
|
||||
});
|
||||
const model = this.selectModel(normalizedCond);
|
||||
|
||||
try {
|
||||
metrics.ai
|
||||
@@ -163,7 +172,11 @@ export class PerplexityProvider extends CopilotProvider<PerplexityConfig> {
|
||||
middleware,
|
||||
});
|
||||
const adapter = this.createNativeAdapter(tools, middleware.node?.text);
|
||||
for await (const chunk of adapter.streamText(request, options.signal)) {
|
||||
for await (const chunk of adapter.streamText(
|
||||
request,
|
||||
options.signal,
|
||||
messages
|
||||
)) {
|
||||
yield chunk;
|
||||
}
|
||||
} catch (e: any) {
|
||||
|
||||
@@ -51,13 +51,21 @@ const DEFAULT_MIDDLEWARE_BY_TYPE: Record<
|
||||
},
|
||||
},
|
||||
[CopilotProviderType.Gemini]: {
|
||||
rust: {
|
||||
request: ['normalize_messages', 'tool_schema_rewrite'],
|
||||
stream: ['stream_event_normalize', 'citation_indexing'],
|
||||
},
|
||||
node: {
|
||||
text: ['callout'],
|
||||
text: ['citation_footnote', 'callout'],
|
||||
},
|
||||
},
|
||||
[CopilotProviderType.GeminiVertex]: {
|
||||
rust: {
|
||||
request: ['normalize_messages', 'tool_schema_rewrite'],
|
||||
stream: ['stream_event_normalize', 'citation_indexing'],
|
||||
},
|
||||
node: {
|
||||
text: ['callout'],
|
||||
text: ['citation_footnote', 'callout'],
|
||||
},
|
||||
},
|
||||
[CopilotProviderType.FAL]: {},
|
||||
|
||||
@@ -5,7 +5,7 @@ import type {
|
||||
ProviderMiddlewareConfig,
|
||||
} from '../config';
|
||||
import { resolveProviderMiddleware } from './provider-middleware';
|
||||
import { CopilotProviderType, type ModelOutputType } from './types';
|
||||
import { CopilotProviderType, ModelOutputType } from './types';
|
||||
|
||||
const PROVIDER_ID_PATTERN = /^[a-zA-Z0-9-_]+$/;
|
||||
|
||||
@@ -239,8 +239,13 @@ export function resolveModel({
|
||||
};
|
||||
}
|
||||
|
||||
const defaultProviderId =
|
||||
outputType && outputType !== ModelOutputType.Rerank
|
||||
? registry.defaults[outputType]
|
||||
: undefined;
|
||||
|
||||
const fallbackOrder = [
|
||||
...(outputType ? [registry.defaults[outputType]] : []),
|
||||
...(defaultProviderId ? [defaultProviderId] : []),
|
||||
registry.defaults.fallback,
|
||||
...registry.order,
|
||||
].filter((id): id is string => !!id);
|
||||
|
||||
@@ -2,7 +2,6 @@ import { AsyncLocalStorage } from 'node:async_hooks';
|
||||
|
||||
import { Inject, Injectable, Logger } from '@nestjs/common';
|
||||
import { ModuleRef } from '@nestjs/core';
|
||||
import { Tool, ToolSet } from 'ai';
|
||||
import { z } from 'zod';
|
||||
|
||||
import {
|
||||
@@ -27,6 +26,8 @@ import {
|
||||
buildDocSearchGetter,
|
||||
buildDocUpdateHandler,
|
||||
buildDocUpdateMetaHandler,
|
||||
type CopilotTool,
|
||||
type CopilotToolSet,
|
||||
createBlobReadTool,
|
||||
createCodeArtifactTool,
|
||||
createConversationSummaryTool,
|
||||
@@ -42,6 +43,7 @@ import {
|
||||
createExaSearchTool,
|
||||
createSectionEditTool,
|
||||
} from '../tools';
|
||||
import { canonicalizePromptAttachment } from './attachments';
|
||||
import { CopilotProviderFactory } from './factory';
|
||||
import { resolveProviderMiddleware } from './provider-middleware';
|
||||
import { buildProviderRegistry } from './provider-registry';
|
||||
@@ -52,12 +54,17 @@ import {
|
||||
type CopilotImageOptions,
|
||||
CopilotProviderModel,
|
||||
CopilotProviderType,
|
||||
type CopilotRerankRequest,
|
||||
CopilotStructuredOptions,
|
||||
EmbeddingMessage,
|
||||
type ModelAttachmentCapability,
|
||||
ModelCapability,
|
||||
ModelConditions,
|
||||
ModelFullConditions,
|
||||
ModelInputType,
|
||||
ModelOutputType,
|
||||
type PromptAttachmentKind,
|
||||
type PromptAttachmentSourceKind,
|
||||
type PromptMessage,
|
||||
PromptMessageSchema,
|
||||
StreamObject,
|
||||
@@ -163,6 +170,163 @@ export abstract class CopilotProvider<C = any> {
|
||||
|
||||
async refreshOnlineModels() {}
|
||||
|
||||
private unique<T>(values: Iterable<T>) {
|
||||
return Array.from(new Set(values));
|
||||
}
|
||||
|
||||
private attachmentKindToInputType(
|
||||
kind: PromptAttachmentKind
|
||||
): ModelInputType {
|
||||
switch (kind) {
|
||||
case 'image':
|
||||
return ModelInputType.Image;
|
||||
case 'audio':
|
||||
return ModelInputType.Audio;
|
||||
default:
|
||||
return ModelInputType.File;
|
||||
}
|
||||
}
|
||||
|
||||
protected async inferModelConditionsFromMessages(
|
||||
messages?: PromptMessage[],
|
||||
withAttachment = true
|
||||
): Promise<Partial<ModelFullConditions>> {
|
||||
if (!messages?.length || !withAttachment) return {};
|
||||
|
||||
const attachmentKinds: PromptAttachmentKind[] = [];
|
||||
const attachmentSourceKinds: PromptAttachmentSourceKind[] = [];
|
||||
const inputTypes: ModelInputType[] = [];
|
||||
let hasRemoteAttachments = false;
|
||||
|
||||
for (const message of messages) {
|
||||
if (!Array.isArray(message.attachments)) continue;
|
||||
|
||||
for (const attachment of message.attachments) {
|
||||
const normalized = await canonicalizePromptAttachment(
|
||||
attachment,
|
||||
message
|
||||
);
|
||||
attachmentKinds.push(normalized.kind);
|
||||
inputTypes.push(this.attachmentKindToInputType(normalized.kind));
|
||||
attachmentSourceKinds.push(normalized.sourceKind);
|
||||
hasRemoteAttachments = hasRemoteAttachments || normalized.isRemote;
|
||||
}
|
||||
}
|
||||
|
||||
return {
|
||||
...(attachmentKinds.length
|
||||
? { attachmentKinds: this.unique(attachmentKinds) }
|
||||
: {}),
|
||||
...(attachmentSourceKinds.length
|
||||
? { attachmentSourceKinds: this.unique(attachmentSourceKinds) }
|
||||
: {}),
|
||||
...(inputTypes.length ? { inputTypes: this.unique(inputTypes) } : {}),
|
||||
...(hasRemoteAttachments ? { hasRemoteAttachments } : {}),
|
||||
};
|
||||
}
|
||||
|
||||
private mergeModelConditions(
|
||||
cond: ModelFullConditions,
|
||||
inferredCond: Partial<ModelFullConditions>
|
||||
): ModelFullConditions {
|
||||
return {
|
||||
...inferredCond,
|
||||
...cond,
|
||||
inputTypes: this.unique([
|
||||
...(inferredCond.inputTypes ?? []),
|
||||
...(cond.inputTypes ?? []),
|
||||
]),
|
||||
attachmentKinds: this.unique([
|
||||
...(inferredCond.attachmentKinds ?? []),
|
||||
...(cond.attachmentKinds ?? []),
|
||||
]),
|
||||
attachmentSourceKinds: this.unique([
|
||||
...(inferredCond.attachmentSourceKinds ?? []),
|
||||
...(cond.attachmentSourceKinds ?? []),
|
||||
]),
|
||||
hasRemoteAttachments:
|
||||
cond.hasRemoteAttachments ?? inferredCond.hasRemoteAttachments,
|
||||
};
|
||||
}
|
||||
|
||||
protected getAttachCapability(
|
||||
model: CopilotProviderModel,
|
||||
outputType: ModelOutputType
|
||||
): ModelAttachmentCapability | undefined {
|
||||
const capability =
|
||||
model.capabilities.find(cap => cap.output.includes(outputType)) ??
|
||||
model.capabilities[0];
|
||||
if (!capability) {
|
||||
return;
|
||||
}
|
||||
return this.resolveAttachmentCapability(capability, outputType);
|
||||
}
|
||||
|
||||
private resolveAttachmentCapability(
|
||||
cap: ModelCapability,
|
||||
outputType?: ModelOutputType
|
||||
): ModelAttachmentCapability | undefined {
|
||||
if (outputType === ModelOutputType.Structured) {
|
||||
return cap.structuredAttachments ?? cap.attachments;
|
||||
}
|
||||
return cap.attachments;
|
||||
}
|
||||
|
||||
private matchesAttachCapability(
|
||||
cap: ModelCapability,
|
||||
cond: ModelFullConditions
|
||||
) {
|
||||
const {
|
||||
attachmentKinds,
|
||||
attachmentSourceKinds,
|
||||
hasRemoteAttachments,
|
||||
outputType,
|
||||
} = cond;
|
||||
|
||||
if (
|
||||
!attachmentKinds?.length &&
|
||||
!attachmentSourceKinds?.length &&
|
||||
!hasRemoteAttachments
|
||||
) {
|
||||
return true;
|
||||
}
|
||||
|
||||
const attachmentCapability = this.resolveAttachmentCapability(
|
||||
cap,
|
||||
outputType
|
||||
);
|
||||
if (!attachmentCapability) {
|
||||
return !attachmentKinds?.some(
|
||||
kind => !cap.input.includes(this.attachmentKindToInputType(kind))
|
||||
);
|
||||
}
|
||||
|
||||
if (
|
||||
attachmentKinds?.some(kind => !attachmentCapability.kinds.includes(kind))
|
||||
) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (
|
||||
attachmentSourceKinds?.length &&
|
||||
attachmentCapability.sourceKinds?.length &&
|
||||
attachmentSourceKinds.some(
|
||||
kind => !attachmentCapability.sourceKinds?.includes(kind)
|
||||
)
|
||||
) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (
|
||||
hasRemoteAttachments &&
|
||||
attachmentCapability.allowRemoteUrls === false
|
||||
) {
|
||||
return false;
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
private findValidModel(
|
||||
cond: ModelFullConditions
|
||||
): CopilotProviderModel | undefined {
|
||||
@@ -170,7 +334,8 @@ export abstract class CopilotProvider<C = any> {
|
||||
const matcher = (cap: ModelCapability) =>
|
||||
(!outputType || cap.output.includes(outputType)) &&
|
||||
(!inputTypes?.length ||
|
||||
inputTypes.every(type => cap.input.includes(type)));
|
||||
inputTypes.every(type => cap.input.includes(type))) &&
|
||||
this.matchesAttachCapability(cap, cond);
|
||||
|
||||
if (modelId) {
|
||||
const hasOnlineModel = this.onlineModelList.includes(modelId);
|
||||
@@ -213,7 +378,7 @@ export abstract class CopilotProvider<C = any> {
|
||||
protected getProviderSpecificTools(
|
||||
_toolName: CopilotChatTools,
|
||||
_model: string
|
||||
): [string, Tool?] | undefined {
|
||||
): [string, CopilotTool?] | undefined {
|
||||
return;
|
||||
}
|
||||
|
||||
@@ -221,8 +386,8 @@ export abstract class CopilotProvider<C = any> {
|
||||
protected async getTools(
|
||||
options: CopilotChatOptions,
|
||||
model: string
|
||||
): Promise<ToolSet> {
|
||||
const tools: ToolSet = {};
|
||||
): Promise<CopilotToolSet> {
|
||||
const tools: CopilotToolSet = {};
|
||||
if (options?.tools?.length) {
|
||||
this.logger.debug(`getTools: ${JSON.stringify(options.tools)}`);
|
||||
const ac = this.moduleRef.get(AccessController, { strict: false });
|
||||
@@ -377,19 +542,14 @@ export abstract class CopilotProvider<C = any> {
|
||||
messages,
|
||||
embeddings,
|
||||
options = {},
|
||||
withAttachment = true,
|
||||
}: {
|
||||
cond: ModelFullConditions;
|
||||
messages?: PromptMessage[];
|
||||
embeddings?: string[];
|
||||
options?: CopilotChatOptions;
|
||||
}) {
|
||||
const model = this.selectModel(cond);
|
||||
const multimodal = model.capabilities.some(c =>
|
||||
[ModelInputType.Image, ModelInputType.Audio].some(t =>
|
||||
c.input.includes(t)
|
||||
)
|
||||
);
|
||||
|
||||
options?: CopilotChatOptions | CopilotStructuredOptions;
|
||||
withAttachment?: boolean;
|
||||
}): Promise<ModelFullConditions> {
|
||||
if (messages) {
|
||||
const { requireContent = true, requireAttachment = false } = options;
|
||||
|
||||
@@ -402,20 +562,56 @@ export abstract class CopilotProvider<C = any> {
|
||||
})
|
||||
.passthrough()
|
||||
.catchall(z.union([z.string(), z.number(), z.date(), z.null()]))
|
||||
.refine(
|
||||
m =>
|
||||
!(multimodal && requireAttachment && m.role === 'user') ||
|
||||
(m.attachments ? m.attachments.length > 0 : true),
|
||||
{ message: 'attachments required in multimodal mode' }
|
||||
)
|
||||
)
|
||||
.optional();
|
||||
|
||||
this.handleZodError(MessageSchema.safeParse(messages));
|
||||
|
||||
const inferredCond = await this.inferModelConditionsFromMessages(
|
||||
messages,
|
||||
withAttachment
|
||||
);
|
||||
const mergedCond = this.mergeModelConditions(cond, inferredCond);
|
||||
const model = this.selectModel(mergedCond);
|
||||
const multimodal = model.capabilities.some(c =>
|
||||
[ModelInputType.Image, ModelInputType.Audio, ModelInputType.File].some(
|
||||
t => c.input.includes(t)
|
||||
)
|
||||
);
|
||||
|
||||
if (
|
||||
multimodal &&
|
||||
requireAttachment &&
|
||||
!messages.some(
|
||||
message =>
|
||||
message.role === 'user' &&
|
||||
Array.isArray(message.attachments) &&
|
||||
message.attachments.length > 0
|
||||
)
|
||||
) {
|
||||
throw new CopilotPromptInvalid(
|
||||
'attachments required in multimodal mode'
|
||||
);
|
||||
}
|
||||
|
||||
if (embeddings) {
|
||||
this.handleZodError(EmbeddingMessage.safeParse(embeddings));
|
||||
}
|
||||
|
||||
return mergedCond;
|
||||
}
|
||||
|
||||
const inferredCond = await this.inferModelConditionsFromMessages(
|
||||
messages,
|
||||
withAttachment
|
||||
);
|
||||
const mergedCond = this.mergeModelConditions(cond, inferredCond);
|
||||
|
||||
if (embeddings) {
|
||||
this.handleZodError(EmbeddingMessage.safeParse(embeddings));
|
||||
}
|
||||
|
||||
return mergedCond;
|
||||
}
|
||||
|
||||
abstract text(
|
||||
@@ -476,7 +672,7 @@ export abstract class CopilotProvider<C = any> {
|
||||
|
||||
async rerank(
|
||||
_model: ModelConditions,
|
||||
_messages: PromptMessage[][],
|
||||
_request: CopilotRerankRequest,
|
||||
_options?: CopilotChatOptions
|
||||
): Promise<number[]> {
|
||||
throw new CopilotProviderNotSupported({
|
||||
|
||||
@@ -1,23 +0,0 @@
|
||||
const GPT_4_RERANK_MODELS = /^(gpt-4(?:$|[.-]))/;
|
||||
const GPT_5_RERANK_LOGPROBS_MODELS = /^(gpt-5\.2(?:$|-))/;
|
||||
|
||||
export const DEFAULT_RERANK_MODEL = 'gpt-5.2';
|
||||
export const OPENAI_RERANK_TOP_LOGPROBS_LIMIT = 5;
|
||||
export const OPENAI_RERANK_MAX_COMPLETION_TOKENS = 16;
|
||||
|
||||
export function supportsRerankModel(model: string): boolean {
|
||||
return (
|
||||
GPT_4_RERANK_MODELS.test(model) || GPT_5_RERANK_LOGPROBS_MODELS.test(model)
|
||||
);
|
||||
}
|
||||
|
||||
export function usesRerankReasoning(model: string): boolean {
|
||||
return GPT_5_RERANK_LOGPROBS_MODELS.test(model);
|
||||
}
|
||||
|
||||
export function normalizeRerankModel(model?: string | null): string {
|
||||
if (model && supportsRerankModel(model)) {
|
||||
return model;
|
||||
}
|
||||
return DEFAULT_RERANK_MODEL;
|
||||
}
|
||||
@@ -124,14 +124,97 @@ export const ChatMessageRole = Object.values(AiPromptRole) as [
|
||||
'user',
|
||||
];
|
||||
|
||||
const AttachmentUrlSchema = z.string().refine(value => {
|
||||
if (value.startsWith('data:')) {
|
||||
return true;
|
||||
}
|
||||
|
||||
try {
|
||||
const url = new URL(value);
|
||||
return (
|
||||
url.protocol === 'http:' ||
|
||||
url.protocol === 'https:' ||
|
||||
url.protocol === 'gs:'
|
||||
);
|
||||
} catch {
|
||||
return false;
|
||||
}
|
||||
}, 'attachments must use https?://, gs:// or data: urls');
|
||||
|
||||
export const PromptAttachmentSourceKindSchema = z.enum([
|
||||
'url',
|
||||
'data',
|
||||
'bytes',
|
||||
'file_handle',
|
||||
]);
|
||||
|
||||
export const PromptAttachmentKindSchema = z.enum(['image', 'audio', 'file']);
|
||||
|
||||
const AttachmentProviderHintSchema = z
|
||||
.object({
|
||||
provider: z.nativeEnum(CopilotProviderType).optional(),
|
||||
kind: PromptAttachmentKindSchema.optional(),
|
||||
})
|
||||
.strict();
|
||||
|
||||
const PromptAttachmentSchema = z.discriminatedUnion('kind', [
|
||||
z
|
||||
.object({
|
||||
kind: z.literal('url'),
|
||||
url: AttachmentUrlSchema,
|
||||
mimeType: z.string().optional(),
|
||||
fileName: z.string().optional(),
|
||||
providerHint: AttachmentProviderHintSchema.optional(),
|
||||
})
|
||||
.strict(),
|
||||
z
|
||||
.object({
|
||||
kind: z.literal('data'),
|
||||
data: z.string(),
|
||||
mimeType: z.string(),
|
||||
encoding: z.enum(['base64', 'utf8']).optional(),
|
||||
fileName: z.string().optional(),
|
||||
providerHint: AttachmentProviderHintSchema.optional(),
|
||||
})
|
||||
.strict(),
|
||||
z
|
||||
.object({
|
||||
kind: z.literal('bytes'),
|
||||
data: z.string(),
|
||||
mimeType: z.string(),
|
||||
encoding: z.literal('base64').optional(),
|
||||
fileName: z.string().optional(),
|
||||
providerHint: AttachmentProviderHintSchema.optional(),
|
||||
})
|
||||
.strict(),
|
||||
z
|
||||
.object({
|
||||
kind: z.literal('file_handle'),
|
||||
fileHandle: z.string().trim().min(1),
|
||||
mimeType: z.string().optional(),
|
||||
fileName: z.string().optional(),
|
||||
providerHint: AttachmentProviderHintSchema.optional(),
|
||||
})
|
||||
.strict(),
|
||||
]);
|
||||
|
||||
export const ChatMessageAttachment = z.union([
|
||||
z.string().url(),
|
||||
AttachmentUrlSchema,
|
||||
z.object({
|
||||
attachment: z.string(),
|
||||
attachment: AttachmentUrlSchema,
|
||||
mimeType: z.string(),
|
||||
}),
|
||||
PromptAttachmentSchema,
|
||||
]);
|
||||
|
||||
export const PromptResponseFormatSchema = z
|
||||
.object({
|
||||
type: z.literal('json_schema'),
|
||||
schema: z.any(),
|
||||
strict: z.boolean().optional(),
|
||||
})
|
||||
.strict();
|
||||
|
||||
export const StreamObjectSchema = z.discriminatedUnion('type', [
|
||||
z.object({
|
||||
type: z.literal('text-delta'),
|
||||
@@ -161,6 +244,7 @@ export const PureMessageSchema = z.object({
|
||||
streamObjects: z.array(StreamObjectSchema).optional().nullable(),
|
||||
attachments: z.array(ChatMessageAttachment).optional().nullable(),
|
||||
params: z.record(z.any()).optional().nullable(),
|
||||
responseFormat: PromptResponseFormatSchema.optional().nullable(),
|
||||
});
|
||||
|
||||
export const PromptMessageSchema = PureMessageSchema.extend({
|
||||
@@ -169,6 +253,12 @@ export const PromptMessageSchema = PureMessageSchema.extend({
|
||||
export type PromptMessage = z.infer<typeof PromptMessageSchema>;
|
||||
export type PromptParams = NonNullable<PromptMessage['params']>;
|
||||
export type StreamObject = z.infer<typeof StreamObjectSchema>;
|
||||
export type PromptAttachment = z.infer<typeof ChatMessageAttachment>;
|
||||
export type PromptAttachmentSourceKind = z.infer<
|
||||
typeof PromptAttachmentSourceKindSchema
|
||||
>;
|
||||
export type PromptAttachmentKind = z.infer<typeof PromptAttachmentKindSchema>;
|
||||
export type PromptResponseFormat = z.infer<typeof PromptResponseFormatSchema>;
|
||||
|
||||
// ========== options ==========
|
||||
|
||||
@@ -194,7 +284,9 @@ export type CopilotChatTools = NonNullable<
|
||||
>[number];
|
||||
|
||||
export const CopilotStructuredOptionsSchema =
|
||||
CopilotProviderOptionsSchema.merge(PromptConfigStrictSchema).optional();
|
||||
CopilotProviderOptionsSchema.merge(PromptConfigStrictSchema)
|
||||
.extend({ schema: z.any().optional(), strict: z.boolean().optional() })
|
||||
.optional();
|
||||
|
||||
export type CopilotStructuredOptions = z.infer<
|
||||
typeof CopilotStructuredOptionsSchema
|
||||
@@ -220,10 +312,22 @@ export type CopilotEmbeddingOptions = z.infer<
|
||||
typeof CopilotEmbeddingOptionsSchema
|
||||
>;
|
||||
|
||||
export type CopilotRerankCandidate = {
|
||||
id?: string;
|
||||
text: string;
|
||||
};
|
||||
|
||||
export type CopilotRerankRequest = {
|
||||
query: string;
|
||||
candidates: CopilotRerankCandidate[];
|
||||
topK?: number;
|
||||
};
|
||||
|
||||
export enum ModelInputType {
|
||||
Text = 'text',
|
||||
Image = 'image',
|
||||
Audio = 'audio',
|
||||
File = 'file',
|
||||
}
|
||||
|
||||
export enum ModelOutputType {
|
||||
@@ -231,12 +335,21 @@ export enum ModelOutputType {
|
||||
Object = 'object',
|
||||
Embedding = 'embedding',
|
||||
Image = 'image',
|
||||
Rerank = 'rerank',
|
||||
Structured = 'structured',
|
||||
}
|
||||
|
||||
export interface ModelAttachmentCapability {
|
||||
kinds: PromptAttachmentKind[];
|
||||
sourceKinds?: PromptAttachmentSourceKind[];
|
||||
allowRemoteUrls?: boolean;
|
||||
}
|
||||
|
||||
export interface ModelCapability {
|
||||
input: ModelInputType[];
|
||||
output: ModelOutputType[];
|
||||
attachments?: ModelAttachmentCapability;
|
||||
structuredAttachments?: ModelAttachmentCapability;
|
||||
defaultForOutputType?: boolean;
|
||||
}
|
||||
|
||||
@@ -248,6 +361,9 @@ export interface CopilotProviderModel {
|
||||
|
||||
export type ModelConditions = {
|
||||
inputTypes?: ModelInputType[];
|
||||
attachmentKinds?: PromptAttachmentKind[];
|
||||
attachmentSourceKinds?: PromptAttachmentSourceKind[];
|
||||
hasRemoteAttachments?: boolean;
|
||||
modelId?: string;
|
||||
};
|
||||
|
||||
|
||||
@@ -1,34 +1,39 @@
|
||||
import { GoogleVertexProviderSettings } from '@ai-sdk/google-vertex';
|
||||
import { GoogleVertexAnthropicProviderSettings } from '@ai-sdk/google-vertex/anthropic';
|
||||
import { Logger } from '@nestjs/common';
|
||||
import {
|
||||
AssistantModelMessage,
|
||||
FilePart,
|
||||
ImagePart,
|
||||
TextPart,
|
||||
TextStreamPart,
|
||||
UserModelMessage,
|
||||
} from 'ai';
|
||||
import { GoogleAuth, GoogleAuthOptions } from 'google-auth-library';
|
||||
import z, { ZodType } from 'zod';
|
||||
import z from 'zod';
|
||||
|
||||
import {
|
||||
bufferToArrayBuffer,
|
||||
fetchBuffer,
|
||||
OneMinute,
|
||||
ResponseTooLargeError,
|
||||
safeFetch,
|
||||
SsrfBlockedError,
|
||||
} from '../../../base';
|
||||
import { CustomAITools } from '../tools';
|
||||
import { PromptMessage, StreamObject } from './types';
|
||||
import { OneMinute, safeFetch } from '../../../base';
|
||||
import { PromptAttachment, StreamObject } from './types';
|
||||
|
||||
type ChatMessage = UserModelMessage | AssistantModelMessage;
|
||||
export type VertexProviderConfig = {
|
||||
location?: string;
|
||||
project?: string;
|
||||
baseURL?: string;
|
||||
googleAuthOptions?: GoogleAuthOptions;
|
||||
fetch?: typeof fetch;
|
||||
};
|
||||
|
||||
export type VertexAnthropicProviderConfig = VertexProviderConfig;
|
||||
|
||||
type CopilotTextStreamPart =
|
||||
| { type: 'text-delta'; text: string; id?: string }
|
||||
| { type: 'reasoning-delta'; text: string; id?: string }
|
||||
| {
|
||||
type: 'tool-call';
|
||||
toolCallId: string;
|
||||
toolName: string;
|
||||
input: Record<string, unknown>;
|
||||
}
|
||||
| {
|
||||
type: 'tool-result';
|
||||
toolCallId: string;
|
||||
toolName: string;
|
||||
input: Record<string, unknown>;
|
||||
output: unknown;
|
||||
}
|
||||
| { type: 'error'; error: unknown };
|
||||
|
||||
const ATTACHMENT_MAX_BYTES = 20 * 1024 * 1024;
|
||||
const ATTACH_HEAD_PARAMS = { timeoutMs: OneMinute / 12, maxRedirects: 3 };
|
||||
|
||||
const SIMPLE_IMAGE_URL_REGEX = /^(https?:\/\/|data:image\/)/;
|
||||
const FORMAT_INFER_MAP: Record<string, string> = {
|
||||
pdf: 'application/pdf',
|
||||
mp3: 'audio/mpeg',
|
||||
@@ -53,9 +58,39 @@ const FORMAT_INFER_MAP: Record<string, string> = {
|
||||
flv: 'video/flv',
|
||||
};
|
||||
|
||||
async function fetchArrayBuffer(url: string): Promise<ArrayBuffer> {
|
||||
const { buffer } = await fetchBuffer(url, ATTACHMENT_MAX_BYTES);
|
||||
return bufferToArrayBuffer(buffer);
|
||||
function toBase64Data(data: string, encoding: 'base64' | 'utf8' = 'base64') {
|
||||
return encoding === 'base64'
|
||||
? data
|
||||
: Buffer.from(data, 'utf8').toString('base64');
|
||||
}
|
||||
|
||||
export function promptAttachmentToUrl(
|
||||
attachment: PromptAttachment
|
||||
): string | undefined {
|
||||
if (typeof attachment === 'string') return attachment;
|
||||
if ('attachment' in attachment) return attachment.attachment;
|
||||
switch (attachment.kind) {
|
||||
case 'url':
|
||||
return attachment.url;
|
||||
case 'data':
|
||||
return `data:${attachment.mimeType};base64,${toBase64Data(
|
||||
attachment.data,
|
||||
attachment.encoding
|
||||
)}`;
|
||||
case 'bytes':
|
||||
return `data:${attachment.mimeType};base64,${attachment.data}`;
|
||||
case 'file_handle':
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
||||
export function promptAttachmentMimeType(
|
||||
attachment: PromptAttachment,
|
||||
fallbackMimeType?: string
|
||||
): string | undefined {
|
||||
if (typeof attachment === 'string') return fallbackMimeType;
|
||||
if ('attachment' in attachment) return attachment.mimeType;
|
||||
return attachment.mimeType ?? fallbackMimeType;
|
||||
}
|
||||
|
||||
export async function inferMimeType(url: string) {
|
||||
@@ -69,346 +104,21 @@ export async function inferMimeType(url: string) {
|
||||
if (ext) {
|
||||
return ext;
|
||||
}
|
||||
try {
|
||||
const mimeType = await safeFetch(
|
||||
url,
|
||||
{ method: 'HEAD' },
|
||||
ATTACH_HEAD_PARAMS
|
||||
).then(res => res.headers.get('content-type'));
|
||||
if (mimeType) return mimeType;
|
||||
} catch {
|
||||
// ignore and fallback to default
|
||||
}
|
||||
}
|
||||
try {
|
||||
const mimeType = await safeFetch(
|
||||
url,
|
||||
{ method: 'HEAD' },
|
||||
ATTACH_HEAD_PARAMS
|
||||
).then(res => res.headers.get('content-type'));
|
||||
if (mimeType) return mimeType;
|
||||
} catch {
|
||||
// ignore and fallback to default
|
||||
}
|
||||
return 'application/octet-stream';
|
||||
}
|
||||
|
||||
export async function chatToGPTMessage(
|
||||
messages: PromptMessage[],
|
||||
// TODO(@darkskygit): move this logic in interface refactoring
|
||||
withAttachment: boolean = true,
|
||||
// NOTE: some providers in vercel ai sdk are not able to handle url attachments yet
|
||||
// so we need to use base64 encoded attachments instead
|
||||
useBase64Attachment: boolean = false
|
||||
): Promise<[string | undefined, ChatMessage[], ZodType?]> {
|
||||
const hasSystem = messages[0]?.role === 'system';
|
||||
const system = hasSystem ? messages[0] : undefined;
|
||||
const normalizedMessages = hasSystem ? messages.slice(1) : messages;
|
||||
const schema =
|
||||
system?.params?.schema && system.params.schema instanceof ZodType
|
||||
? system.params.schema
|
||||
: undefined;
|
||||
|
||||
// filter redundant fields
|
||||
const msgs: ChatMessage[] = [];
|
||||
for (let { role, content, attachments, params } of normalizedMessages.filter(
|
||||
m => m.role !== 'system'
|
||||
)) {
|
||||
content = content.trim();
|
||||
role = role as 'user' | 'assistant';
|
||||
const mimetype = params?.mimetype;
|
||||
if (Array.isArray(attachments)) {
|
||||
const contents: (TextPart | ImagePart | FilePart)[] = [];
|
||||
if (content.length) {
|
||||
contents.push({ type: 'text', text: content });
|
||||
}
|
||||
|
||||
if (withAttachment) {
|
||||
for (let attachment of attachments) {
|
||||
let mediaType: string;
|
||||
if (typeof attachment === 'string') {
|
||||
mediaType =
|
||||
typeof mimetype === 'string'
|
||||
? mimetype
|
||||
: await inferMimeType(attachment);
|
||||
} else {
|
||||
({ attachment, mimeType: mediaType } = attachment);
|
||||
}
|
||||
if (SIMPLE_IMAGE_URL_REGEX.test(attachment)) {
|
||||
const data =
|
||||
attachment.startsWith('data:') || useBase64Attachment
|
||||
? await fetchArrayBuffer(attachment).catch(error => {
|
||||
// Avoid leaking internal details for blocked URLs.
|
||||
if (
|
||||
error instanceof SsrfBlockedError ||
|
||||
error instanceof ResponseTooLargeError
|
||||
) {
|
||||
throw new Error('Attachment URL is not allowed');
|
||||
}
|
||||
throw error;
|
||||
})
|
||||
: new URL(attachment);
|
||||
if (mediaType.startsWith('image/')) {
|
||||
contents.push({ type: 'image', image: data, mediaType });
|
||||
} else {
|
||||
contents.push({ type: 'file' as const, data, mediaType });
|
||||
}
|
||||
}
|
||||
}
|
||||
} else if (!content.length) {
|
||||
// temp fix for pplx
|
||||
contents.push({ type: 'text', text: '[no content]' });
|
||||
}
|
||||
|
||||
msgs.push({ role, content: contents } as ChatMessage);
|
||||
} else {
|
||||
msgs.push({ role, content });
|
||||
}
|
||||
}
|
||||
|
||||
return [system?.content, msgs, schema];
|
||||
}
|
||||
|
||||
// pattern types the callback will receive
|
||||
type Pattern =
|
||||
| { kind: 'index'; value: number } // [123]
|
||||
| { kind: 'link'; text: string; url: string } // [text](url)
|
||||
| { kind: 'wrappedLink'; text: string; url: string }; // ([text](url))
|
||||
|
||||
type NeedMore = { kind: 'needMore' };
|
||||
type Failed = { kind: 'fail'; nextPos: number };
|
||||
type Finished =
|
||||
| { kind: 'ok'; endPos: number; text: string; url: string }
|
||||
| { kind: 'index'; endPos: number; value: number };
|
||||
type ParseStatus = Finished | NeedMore | Failed;
|
||||
|
||||
type PatternCallback = (m: Pattern) => string;
|
||||
|
||||
export class StreamPatternParser {
|
||||
#buffer = '';
|
||||
|
||||
constructor(private readonly callback: PatternCallback) {}
|
||||
|
||||
write(chunk: string): string {
|
||||
this.#buffer += chunk;
|
||||
const output: string[] = [];
|
||||
let i = 0;
|
||||
|
||||
while (i < this.#buffer.length) {
|
||||
const ch = this.#buffer[i];
|
||||
|
||||
// [[[number]]] or [text](url) or ([text](url))
|
||||
if (ch === '[' || (ch === '(' && this.peek(i + 1) === '[')) {
|
||||
const isWrapped = ch === '(';
|
||||
const startPos = isWrapped ? i + 1 : i;
|
||||
const res = this.tryParse(startPos);
|
||||
if (res.kind === 'needMore') break;
|
||||
const { output: out, nextPos } = this.handlePattern(
|
||||
res,
|
||||
isWrapped,
|
||||
startPos,
|
||||
i
|
||||
);
|
||||
output.push(out);
|
||||
i = nextPos;
|
||||
continue;
|
||||
}
|
||||
output.push(ch);
|
||||
i += 1;
|
||||
}
|
||||
|
||||
this.#buffer = this.#buffer.slice(i);
|
||||
return output.join('');
|
||||
}
|
||||
|
||||
end(): string {
|
||||
const rest = this.#buffer;
|
||||
this.#buffer = '';
|
||||
return rest;
|
||||
}
|
||||
|
||||
// =========== helpers ===========
|
||||
|
||||
private peek(pos: number): string | undefined {
|
||||
return pos < this.#buffer.length ? this.#buffer[pos] : undefined;
|
||||
}
|
||||
|
||||
private tryParse(pos: number): ParseStatus {
|
||||
const nestedRes = this.tryParseNestedIndex(pos);
|
||||
if (nestedRes) return nestedRes;
|
||||
return this.tryParseBracketPattern(pos);
|
||||
}
|
||||
|
||||
private tryParseNestedIndex(pos: number): ParseStatus | null {
|
||||
if (this.peek(pos + 1) !== '[') return null;
|
||||
|
||||
let i = pos;
|
||||
let bracketCount = 0;
|
||||
|
||||
while (i < this.#buffer.length && this.#buffer[i] === '[') {
|
||||
bracketCount++;
|
||||
i++;
|
||||
}
|
||||
|
||||
if (bracketCount >= 2) {
|
||||
if (i >= this.#buffer.length) {
|
||||
return { kind: 'needMore' };
|
||||
}
|
||||
|
||||
let content = '';
|
||||
while (i < this.#buffer.length && this.#buffer[i] !== ']') {
|
||||
content += this.#buffer[i++];
|
||||
}
|
||||
|
||||
let rightBracketCount = 0;
|
||||
while (i < this.#buffer.length && this.#buffer[i] === ']') {
|
||||
rightBracketCount++;
|
||||
i++;
|
||||
}
|
||||
|
||||
if (i >= this.#buffer.length && rightBracketCount < bracketCount) {
|
||||
return { kind: 'needMore' };
|
||||
}
|
||||
|
||||
if (
|
||||
rightBracketCount === bracketCount &&
|
||||
content.length > 0 &&
|
||||
this.isNumeric(content)
|
||||
) {
|
||||
if (this.peek(i) === '(') {
|
||||
return { kind: 'fail', nextPos: i };
|
||||
}
|
||||
return { kind: 'index', endPos: i, value: Number(content) };
|
||||
}
|
||||
}
|
||||
|
||||
return null;
|
||||
}
|
||||
|
||||
private tryParseBracketPattern(pos: number): ParseStatus {
|
||||
let i = pos + 1; // skip '['
|
||||
if (i >= this.#buffer.length) {
|
||||
return { kind: 'needMore' };
|
||||
}
|
||||
|
||||
let content = '';
|
||||
while (i < this.#buffer.length && this.#buffer[i] !== ']') {
|
||||
const nextChar = this.#buffer[i];
|
||||
if (nextChar === '[') {
|
||||
return { kind: 'fail', nextPos: i };
|
||||
}
|
||||
content += nextChar;
|
||||
i += 1;
|
||||
}
|
||||
|
||||
if (i >= this.#buffer.length) {
|
||||
return { kind: 'needMore' };
|
||||
}
|
||||
const after = i + 1;
|
||||
const afterChar = this.peek(after);
|
||||
|
||||
if (content.length > 0 && this.isNumeric(content) && afterChar !== '(') {
|
||||
// [number] pattern
|
||||
return { kind: 'index', endPos: after, value: Number(content) };
|
||||
} else if (afterChar !== '(') {
|
||||
// [text](url) pattern
|
||||
return { kind: 'fail', nextPos: after };
|
||||
}
|
||||
|
||||
i = after + 1; // skip '('
|
||||
if (i >= this.#buffer.length) {
|
||||
return { kind: 'needMore' };
|
||||
}
|
||||
|
||||
let url = '';
|
||||
while (i < this.#buffer.length && this.#buffer[i] !== ')') {
|
||||
url += this.#buffer[i++];
|
||||
}
|
||||
if (i >= this.#buffer.length) {
|
||||
return { kind: 'needMore' };
|
||||
}
|
||||
return { kind: 'ok', endPos: i + 1, text: content, url };
|
||||
}
|
||||
|
||||
private isNumeric(str: string): boolean {
|
||||
return !Number.isNaN(Number(str)) && str.trim() !== '';
|
||||
}
|
||||
|
||||
private handlePattern(
|
||||
pattern: Finished | Failed,
|
||||
isWrapped: boolean,
|
||||
start: number,
|
||||
current: number
|
||||
): { output: string; nextPos: number } {
|
||||
if (pattern.kind === 'fail') {
|
||||
return {
|
||||
output: this.#buffer.slice(current, pattern.nextPos),
|
||||
nextPos: pattern.nextPos,
|
||||
};
|
||||
}
|
||||
|
||||
if (isWrapped) {
|
||||
const afterLinkPos = pattern.endPos;
|
||||
if (this.peek(afterLinkPos) !== ')') {
|
||||
if (afterLinkPos >= this.#buffer.length) {
|
||||
return { output: '', nextPos: current };
|
||||
}
|
||||
return { output: '(', nextPos: start };
|
||||
}
|
||||
|
||||
const out =
|
||||
pattern.kind === 'index'
|
||||
? this.callback({ ...pattern, kind: 'index' })
|
||||
: this.callback({ ...pattern, kind: 'wrappedLink' });
|
||||
return { output: out, nextPos: afterLinkPos + 1 };
|
||||
} else {
|
||||
const out =
|
||||
pattern.kind === 'ok'
|
||||
? this.callback({ ...pattern, kind: 'link' })
|
||||
: this.callback({ ...pattern, kind: 'index' });
|
||||
return { output: out, nextPos: pattern.endPos };
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
export class CitationParser {
|
||||
private readonly citations: string[] = [];
|
||||
|
||||
private readonly parser = new StreamPatternParser(p => {
|
||||
switch (p.kind) {
|
||||
case 'index': {
|
||||
if (p.value <= this.citations.length) {
|
||||
return `[^${p.value}]`;
|
||||
}
|
||||
return `[${p.value}]`;
|
||||
}
|
||||
case 'wrappedLink': {
|
||||
const index = this.citations.indexOf(p.url);
|
||||
if (index === -1) {
|
||||
this.citations.push(p.url);
|
||||
return `[^${this.citations.length}]`;
|
||||
}
|
||||
return `[^${index + 1}]`;
|
||||
}
|
||||
case 'link': {
|
||||
return `[${p.text}](${p.url})`;
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
public push(citation: string) {
|
||||
this.citations.push(citation);
|
||||
}
|
||||
|
||||
public parse(content: string) {
|
||||
return this.parser.write(content);
|
||||
}
|
||||
|
||||
public end() {
|
||||
return this.parser.end() + '\n' + this.getFootnotes();
|
||||
}
|
||||
|
||||
private getFootnotes() {
|
||||
const footnotes = this.citations.map((citation, index) => {
|
||||
return `[^${index + 1}]: {"type":"url","url":"${encodeURIComponent(
|
||||
citation
|
||||
)}"}`;
|
||||
});
|
||||
return footnotes.join('\n');
|
||||
}
|
||||
}
|
||||
|
||||
export type CitationIndexedEvent = {
|
||||
type CitationIndexedEvent = {
|
||||
type: 'citation';
|
||||
index: number;
|
||||
url: string;
|
||||
@@ -436,7 +146,7 @@ export class CitationFootnoteFormatter {
|
||||
}
|
||||
}
|
||||
|
||||
type ChunkType = TextStreamPart<CustomAITools>['type'];
|
||||
type ChunkType = CopilotTextStreamPart['type'];
|
||||
|
||||
export function toError(error: unknown): Error {
|
||||
if (typeof error === 'string') {
|
||||
@@ -458,6 +168,14 @@ type DocEditFootnote = {
|
||||
intent: string;
|
||||
result: string;
|
||||
};
|
||||
|
||||
function asRecord(value: unknown): Record<string, unknown> | null {
|
||||
if (value && typeof value === 'object' && !Array.isArray(value)) {
|
||||
return value as Record<string, unknown>;
|
||||
}
|
||||
return null;
|
||||
}
|
||||
|
||||
export class TextStreamParser {
|
||||
private readonly logger = new Logger(TextStreamParser.name);
|
||||
private readonly CALLOUT_PREFIX = '\n[!]\n';
|
||||
@@ -468,7 +186,7 @@ export class TextStreamParser {
|
||||
|
||||
private readonly docEditFootnotes: DocEditFootnote[] = [];
|
||||
|
||||
public parse(chunk: TextStreamPart<CustomAITools>) {
|
||||
public parse(chunk: CopilotTextStreamPart) {
|
||||
let result = '';
|
||||
switch (chunk.type) {
|
||||
case 'text-delta': {
|
||||
@@ -517,7 +235,7 @@ export class TextStreamParser {
|
||||
}
|
||||
case 'doc_edit': {
|
||||
this.docEditFootnotes.push({
|
||||
intent: chunk.input.instructions,
|
||||
intent: String(chunk.input.instructions ?? ''),
|
||||
result: '',
|
||||
});
|
||||
break;
|
||||
@@ -533,14 +251,12 @@ export class TextStreamParser {
|
||||
result = this.addPrefix(result);
|
||||
switch (chunk.toolName) {
|
||||
case 'doc_edit': {
|
||||
const array =
|
||||
chunk.output && typeof chunk.output === 'object'
|
||||
? chunk.output.result
|
||||
: undefined;
|
||||
const output = asRecord(chunk.output);
|
||||
const array = output?.result;
|
||||
if (Array.isArray(array)) {
|
||||
result += array
|
||||
.map(item => {
|
||||
return `\n${item.changedContent}\n`;
|
||||
return `\n${String(asRecord(item)?.changedContent ?? '')}\n`;
|
||||
})
|
||||
.join('');
|
||||
this.docEditFootnotes[this.docEditFootnotes.length - 1].result =
|
||||
@@ -557,8 +273,11 @@ export class TextStreamParser {
|
||||
} else if (typeof output === 'string') {
|
||||
result += `\n${output}\n`;
|
||||
} else {
|
||||
const message = asRecord(output)?.message;
|
||||
this.logger.warn(
|
||||
`Unexpected result type for doc_semantic_search: ${output?.message || 'Unknown error'}`
|
||||
`Unexpected result type for doc_semantic_search: ${
|
||||
typeof message === 'string' ? message : 'Unknown error'
|
||||
}`
|
||||
);
|
||||
}
|
||||
break;
|
||||
@@ -572,9 +291,11 @@ export class TextStreamParser {
|
||||
break;
|
||||
}
|
||||
case 'doc_compose': {
|
||||
const output = chunk.output;
|
||||
if (output && typeof output === 'object' && 'title' in output) {
|
||||
result += `\nDocument "${output.title}" created successfully with ${output.wordCount} words.\n`;
|
||||
const output = asRecord(chunk.output);
|
||||
if (output && typeof output.title === 'string') {
|
||||
result += `\nDocument "${output.title}" created successfully with ${String(
|
||||
output.wordCount ?? 0
|
||||
)} words.\n`;
|
||||
}
|
||||
break;
|
||||
}
|
||||
@@ -654,7 +375,7 @@ export class TextStreamParser {
|
||||
}
|
||||
|
||||
export class StreamObjectParser {
|
||||
public parse(chunk: TextStreamPart<CustomAITools>) {
|
||||
public parse(chunk: CopilotTextStreamPart) {
|
||||
switch (chunk.type) {
|
||||
case 'reasoning-delta': {
|
||||
return { type: 'reasoning' as const, textDelta: chunk.text };
|
||||
@@ -747,9 +468,7 @@ function normalizeUrl(baseURL?: string) {
|
||||
}
|
||||
}
|
||||
|
||||
export function getVertexAnthropicBaseUrl(
|
||||
options: GoogleVertexAnthropicProviderSettings
|
||||
) {
|
||||
export function getVertexAnthropicBaseUrl(options: VertexProviderConfig) {
|
||||
const normalizedBaseUrl = normalizeUrl(options.baseURL);
|
||||
if (normalizedBaseUrl) return normalizedBaseUrl;
|
||||
const { location, project } = options;
|
||||
@@ -758,7 +477,7 @@ export function getVertexAnthropicBaseUrl(
|
||||
}
|
||||
|
||||
export async function getGoogleAuth(
|
||||
options: GoogleVertexAnthropicProviderSettings | GoogleVertexProviderSettings,
|
||||
options: VertexProviderConfig,
|
||||
publisher: 'anthropic' | 'google'
|
||||
) {
|
||||
function getBaseUrl() {
|
||||
@@ -777,7 +496,7 @@ export async function getGoogleAuth(
|
||||
}
|
||||
const auth = new GoogleAuth({
|
||||
scopes: ['https://www.googleapis.com/auth/cloud-platform'],
|
||||
...(options.googleAuthOptions as GoogleAuthOptions),
|
||||
...options.googleAuthOptions,
|
||||
});
|
||||
const client = await auth.getClient();
|
||||
const token = await client.getAccessToken();
|
||||
|
||||
Reference in New Issue
Block a user