feat(server): refactor provider interface (#11665)

fix AI-4
fix AI-18

better provider/model choose to allow fallback to similar models (e.g., self-hosted) when the provider is not fully configured
split functions of different output types
This commit is contained in:
darkskygit
2025-05-22 06:28:20 +00:00
parent a3b8aaff61
commit b388f92c96
36 changed files with 1465 additions and 903 deletions
@@ -21,19 +21,21 @@ import {
} from '../../../base';
import { createExaCrawlTool, createExaSearchTool } from '../tools';
import { CopilotProvider } from './provider';
import {
ChatMessageRole,
CopilotCapability,
import type {
CopilotChatOptions,
CopilotEmbeddingOptions,
CopilotImageOptions,
CopilotImageToTextProvider,
CopilotProviderType,
CopilotTextToEmbeddingProvider,
CopilotTextToImageProvider,
CopilotTextToTextProvider,
CopilotStructuredOptions,
ModelConditions,
ModelFullConditions,
PromptMessage,
} from './types';
import {
ChatMessageRole,
CopilotProviderType,
ModelInputType,
ModelOutputType,
} from './types';
import { chatToGPTMessage, CitationParser } from './utils';
export const DEFAULT_DIMENSIONS = 256;
@@ -49,44 +51,144 @@ type OpenAITools = {
web_crawl_exa: ReturnType<typeof createExaCrawlTool>;
};
export class OpenAIProvider
extends CopilotProvider<OpenAIConfig>
implements
CopilotTextToTextProvider,
CopilotTextToEmbeddingProvider,
CopilotTextToImageProvider,
CopilotImageToTextProvider
{
export class OpenAIProvider extends CopilotProvider<OpenAIConfig> {
readonly type = CopilotProviderType.OpenAI;
readonly capabilities = [
CopilotCapability.TextToText,
CopilotCapability.TextToEmbedding,
CopilotCapability.TextToImage,
CopilotCapability.ImageToText,
];
readonly models = [
// text to text
'gpt-4o',
'gpt-4o-2024-08-06',
'gpt-4o-mini',
'gpt-4o-mini-2024-07-18',
'gpt-4.1',
'gpt-4.1-2025-04-14',
'gpt-4.1-mini',
'o1',
'o3',
'o4-mini',
// embeddings
'text-embedding-3-large',
'text-embedding-3-small',
'text-embedding-ada-002',
// moderation
'text-moderation-latest',
'text-moderation-stable',
// text to image
'dall-e-3',
'gpt-image-1',
// Text to Text models
{
id: 'gpt-4o',
capabilities: [
{
input: [ModelInputType.Text, ModelInputType.Image],
output: [ModelOutputType.Text],
},
],
},
// FIXME(@darkskygit): deprecated
{
id: 'gpt-4o-2024-08-06',
capabilities: [
{
input: [ModelInputType.Text, ModelInputType.Image],
output: [ModelOutputType.Text],
},
],
},
{
id: 'gpt-4o-mini',
capabilities: [
{
input: [ModelInputType.Text, ModelInputType.Image],
output: [ModelOutputType.Text],
},
],
},
// FIXME(@darkskygit): deprecated
{
id: 'gpt-4o-mini-2024-07-18',
capabilities: [
{
input: [ModelInputType.Text, ModelInputType.Image],
output: [ModelOutputType.Text],
},
],
},
{
id: 'gpt-4.1',
capabilities: [
{
input: [ModelInputType.Text, ModelInputType.Image],
output: [ModelOutputType.Text],
defaultForOutputType: true,
},
],
},
{
id: 'gpt-4.1-2025-04-14',
capabilities: [
{
input: [ModelInputType.Text, ModelInputType.Image],
output: [ModelOutputType.Text],
},
],
},
{
id: 'gpt-4.1-mini',
capabilities: [
{
input: [ModelInputType.Text, ModelInputType.Image],
output: [ModelOutputType.Text],
},
],
},
{
id: 'o1',
capabilities: [
{
input: [ModelInputType.Text, ModelInputType.Image],
output: [ModelOutputType.Text],
},
],
},
{
id: 'o3',
capabilities: [
{
input: [ModelInputType.Text, ModelInputType.Image],
output: [ModelOutputType.Text],
},
],
},
{
id: 'o4-mini',
capabilities: [
{
input: [ModelInputType.Text, ModelInputType.Image],
output: [ModelOutputType.Text],
},
],
},
// Embedding models
{
id: 'text-embedding-3-large',
capabilities: [
{
input: [ModelInputType.Text],
output: [ModelOutputType.Embedding],
defaultForOutputType: true,
},
],
},
{
id: 'text-embedding-3-small',
capabilities: [
{
input: [ModelInputType.Text],
output: [ModelOutputType.Embedding],
},
],
},
// Image generation models
{
id: 'dall-e-3',
capabilities: [
{
input: [ModelInputType.Text],
output: [ModelOutputType.Image],
},
],
},
{
id: 'gpt-image-1',
capabilities: [
{
input: [ModelInputType.Text, ModelInputType.Image],
output: [ModelOutputType.Image],
defaultForOutputType: true,
},
],
},
];
private readonly MAX_STEPS = 20;
@@ -108,18 +210,17 @@ export class OpenAIProvider
}
protected async checkParams({
cond,
messages,
embeddings,
model,
options = {},
}: {
cond: ModelFullConditions;
messages?: PromptMessage[];
embeddings?: string[];
model: string;
options: CopilotChatOptions;
options?: CopilotChatOptions;
}) {
if (!(await this.isModelAvailable(model))) {
throw new CopilotPromptInvalid(`Invalid model: ${model}`);
if (!(await this.match(cond))) {
throw new CopilotPromptInvalid(`Invalid model: ${cond.modelId}`);
}
if (Array.isArray(messages) && messages.length > 0) {
if (
@@ -147,14 +248,6 @@ export class OpenAIProvider
) {
throw new CopilotPromptInvalid('Invalid message role');
}
// json mode need 'json' keyword in content
// ref: https://platform.openai.com/docs/api-reference/chat/create#chat-create-response_format
if (
options.jsonMode &&
!messages.some(m => m.content.toLowerCase().includes('json'))
) {
throw new CopilotPromptInvalid('Prompt not support json mode');
}
} else if (
Array.isArray(embeddings) &&
embeddings.some(e => typeof e !== 'string' || !e || !e.trim())
@@ -215,82 +308,77 @@ export class OpenAIProvider
return tools;
}
// ====== text to text ======
async generateText(
async text(
cond: ModelConditions,
messages: PromptMessage[],
model: string = 'gpt-4.1-mini',
options: CopilotChatOptions = {}
): Promise<string> {
await this.checkParams({ messages, model, options });
const fullCond = {
...cond,
outputType: ModelOutputType.Text,
};
await this.checkParams({ messages, cond: fullCond, options });
const model = this.selectModel(fullCond);
try {
metrics.ai.counter('chat_text_calls').add(1, { model });
metrics.ai.counter('chat_text_calls').add(1, { model: model.id });
const [system, msgs, schema] = await chatToGPTMessage(messages);
const [system, msgs] = await chatToGPTMessage(messages);
const modelInstance = this.#instance(model, {
structuredOutputs: Boolean(options.jsonMode),
user: options.user,
});
const modelInstance = this.#instance.responses(model.id);
const commonParams = {
const { text } = await generateText({
model: modelInstance,
system,
messages: msgs,
temperature: options.temperature || 0,
maxTokens: options.maxTokens || 4096,
providerOptions: {
openai: this.getOpenAIOptions(options, model.id),
},
tools: this.getTools(options, model.id),
maxSteps: this.MAX_STEPS,
abortSignal: options.signal,
};
const { text } = schema
? await generateObject({
...commonParams,
schema,
}).then(r => ({ text: JSON.stringify(r.object) }))
: await generateText({
...commonParams,
providerOptions: {
openai: this.getOpenAIOptions(options, model),
},
tools: this.getTools(options, model),
maxSteps: this.MAX_STEPS,
});
});
return text.trim();
} catch (e: any) {
metrics.ai.counter('chat_text_errors').add(1, { model });
throw this.handleError(e, model, options);
metrics.ai.counter('chat_text_errors').add(1, { model: model.id });
throw this.handleError(e, model.id, options);
}
}
async *generateTextStream(
async *streamText(
cond: ModelConditions,
messages: PromptMessage[],
model: string = 'gpt-4.1-mini',
options: CopilotChatOptions = {}
): AsyncIterable<string> {
await this.checkParams({ messages, model, options });
const fullCond = {
...cond,
outputType: ModelOutputType.Text,
};
await this.checkParams({ messages, cond: fullCond });
const model = this.selectModel(fullCond);
try {
metrics.ai.counter('chat_text_stream_calls').add(1, { model });
metrics.ai.counter('chat_text_stream_calls').add(1, { model: model.id });
const [system, msgs] = await chatToGPTMessage(messages);
const modelInstance = this.#instance.responses(model);
const modelInstance = this.#instance.responses(model.id);
const tools = this.getTools(options, model);
const { fullStream } = streamText({
model: modelInstance,
system,
messages: msgs,
providerOptions: {
openai: this.getOpenAIOptions(options, model),
},
tools: tools as OpenAITools,
maxSteps: this.MAX_STEPS,
frequencyPenalty: options.frequencyPenalty || 0,
presencePenalty: options.presencePenalty || 0,
temperature: options.temperature || 0,
maxTokens: options.maxTokens || 4096,
providerOptions: {
openai: this.getOpenAIOptions(options, model.id),
},
tools: this.getTools(options, model.id) as OpenAITools,
maxSteps: this.MAX_STEPS,
abortSignal: options.signal,
});
@@ -368,54 +456,68 @@ export class OpenAIProvider
}
}
} catch (e: any) {
metrics.ai.counter('chat_text_stream_errors').add(1, { model });
throw this.handleError(e, model, options);
metrics.ai.counter('chat_text_stream_errors').add(1, { model: model.id });
throw this.handleError(e, model.id, options);
}
}
// ====== text to embedding ======
async generateEmbedding(
messages: string | string[],
model: string,
options: CopilotEmbeddingOptions = { dimensions: DEFAULT_DIMENSIONS }
): Promise<number[][]> {
messages = Array.isArray(messages) ? messages : [messages];
await this.checkParams({ embeddings: messages, model, options });
override async structure(
cond: ModelConditions,
messages: PromptMessage[],
options: CopilotStructuredOptions = {}
): Promise<string> {
const fullCond = { ...cond, outputType: ModelOutputType.Structured };
await this.checkParams({ messages, cond: fullCond, options });
const model = this.selectModel(fullCond);
try {
metrics.ai.counter('generate_embedding_calls').add(1, { model });
metrics.ai.counter('chat_text_calls').add(1, { model: model.id });
const modelInstance = this.#instance.embedding(model, {
dimensions: options.dimensions || DEFAULT_DIMENSIONS,
user: options.user,
});
const [system, msgs, schema] = await chatToGPTMessage(messages);
if (!schema) {
throw new CopilotPromptInvalid('Schema is required');
}
const { embeddings } = await embedMany({
const modelInstance = this.#instance.responses(model.id);
const { object } = await generateObject({
model: modelInstance,
values: messages,
system,
messages: msgs,
temperature: ('temperature' in options && options.temperature) || 0,
maxTokens: ('maxTokens' in options && options.maxTokens) || 4096,
schema,
providerOptions: {
openai: options.user ? { user: options.user } : {},
},
abortSignal: options.signal,
});
return embeddings.filter(v => v && Array.isArray(v));
return JSON.stringify(object);
} catch (e: any) {
metrics.ai.counter('generate_embedding_errors').add(1, { model });
throw this.handleError(e, model, options);
metrics.ai.counter('chat_text_errors').add(1, { model: model.id });
throw this.handleError(e, model.id, options);
}
}
// ====== text to image ======
async generateImages(
override async *streamImages(
cond: ModelConditions,
messages: PromptMessage[],
model: string = 'dall-e-3',
options: CopilotImageOptions = {}
): Promise<Array<string>> {
const { content: prompt } = messages.pop() || {};
) {
const fullCond = { ...cond, outputType: ModelOutputType.Image };
await this.checkParams({ messages, cond: fullCond });
const model = this.selectModel(fullCond);
metrics.ai
.counter('generate_images_stream_calls')
.add(1, { model: model.id });
const { content: prompt } = [...messages].pop() || {};
if (!prompt) throw new CopilotPromptInvalid('Prompt is required');
try {
metrics.ai.counter('generate_images_calls').add(1, { model });
const modelInstance = this.#instance.image(model);
const modelInstance = this.#instance.image(model.id);
const result = await generateImage({
model: modelInstance,
@@ -427,29 +529,54 @@ export class OpenAIProvider
},
});
return result.images.map(
const imageUrls = result.images.map(
image => `data:image/png;base64,${image.base64}`
);
for (const imageUrl of imageUrls) {
yield imageUrl;
if (options.signal?.aborted) {
break;
}
}
return;
} catch (e: any) {
metrics.ai.counter('generate_images_errors').add(1, { model });
throw this.handleError(e, model, options);
metrics.ai.counter('generate_images_errors').add(1, { model: model.id });
throw this.handleError(e, model.id, options);
}
}
async *generateImagesStream(
messages: PromptMessage[],
model: string = 'dall-e-3',
options: CopilotImageOptions = {}
): AsyncIterable<string> {
override async embedding(
cond: ModelConditions,
messages: string | string[],
options: CopilotEmbeddingOptions = { dimensions: DEFAULT_DIMENSIONS }
): Promise<number[][]> {
messages = Array.isArray(messages) ? messages : [messages];
const fullCond = { ...cond, outputType: ModelOutputType.Embedding };
await this.checkParams({ embeddings: messages, cond: fullCond, options });
const model = this.selectModel(fullCond);
try {
metrics.ai.counter('generate_images_stream_calls').add(1, { model });
const ret = await this.generateImages(messages, model, options);
for (const url of ret) {
yield url;
}
} catch (e) {
metrics.ai.counter('generate_images_stream_errors').add(1, { model });
throw e;
metrics.ai
.counter('generate_embedding_calls')
.add(1, { model: model.id });
const modelInstance = this.#instance.embedding(model.id, {
dimensions: options.dimensions || DEFAULT_DIMENSIONS,
user: options.user,
});
const { embeddings } = await embedMany({
model: modelInstance,
values: messages,
});
return embeddings.filter(v => v && Array.isArray(v));
} catch (e: any) {
metrics.ai
.counter('generate_embedding_errors')
.add(1, { model: model.id });
throw this.handleError(e, model.id, options);
}
}