fix: pick copilot provider depend on model (#6540)
This commit is contained in:
@@ -43,6 +43,12 @@ export const prompts: Prompt[] = [
|
|||||||
model: '110602490-lcm-sd15-i2i',
|
model: '110602490-lcm-sd15-i2i',
|
||||||
messages: [],
|
messages: [],
|
||||||
},
|
},
|
||||||
|
{
|
||||||
|
name: 'debug:action:fal-sdturbo',
|
||||||
|
action: 'image',
|
||||||
|
model: 'fast-turbo-diffusion',
|
||||||
|
messages: [],
|
||||||
|
},
|
||||||
{
|
{
|
||||||
name: 'Summary',
|
name: 'Summary',
|
||||||
action: 'Summary',
|
action: 'Summary',
|
||||||
|
|||||||
@@ -89,8 +89,10 @@ export class CopilotController {
|
|||||||
@Query('messageId') messageId: string | undefined,
|
@Query('messageId') messageId: string | undefined,
|
||||||
@Query() params: Record<string, string | string[]>
|
@Query() params: Record<string, string | string[]>
|
||||||
): Promise<string> {
|
): Promise<string> {
|
||||||
|
const model = await this.chatSession.get(sessionId).then(s => s?.model);
|
||||||
const provider = this.provider.getProviderByCapability(
|
const provider = this.provider.getProviderByCapability(
|
||||||
CopilotCapability.TextToText
|
CopilotCapability.TextToText,
|
||||||
|
model
|
||||||
);
|
);
|
||||||
if (!provider) {
|
if (!provider) {
|
||||||
throw new InternalServerErrorException('No provider available');
|
throw new InternalServerErrorException('No provider available');
|
||||||
@@ -139,8 +141,10 @@ export class CopilotController {
|
|||||||
@Query('messageId') messageId: string | undefined,
|
@Query('messageId') messageId: string | undefined,
|
||||||
@Query() params: Record<string, string>
|
@Query() params: Record<string, string>
|
||||||
): Promise<Observable<ChatEvent>> {
|
): Promise<Observable<ChatEvent>> {
|
||||||
|
const model = await this.chatSession.get(sessionId).then(s => s?.model);
|
||||||
const provider = this.provider.getProviderByCapability(
|
const provider = this.provider.getProviderByCapability(
|
||||||
CopilotCapability.TextToText
|
CopilotCapability.TextToText,
|
||||||
|
model
|
||||||
);
|
);
|
||||||
if (!provider) {
|
if (!provider) {
|
||||||
throw new InternalServerErrorException('No provider available');
|
throw new InternalServerErrorException('No provider available');
|
||||||
@@ -194,10 +198,13 @@ export class CopilotController {
|
|||||||
@Query('messageId') messageId: string | undefined,
|
@Query('messageId') messageId: string | undefined,
|
||||||
@Query() params: Record<string, string>
|
@Query() params: Record<string, string>
|
||||||
): Promise<Observable<ChatEvent>> {
|
): Promise<Observable<ChatEvent>> {
|
||||||
|
const hasAttachment = await this.hasAttachment(sessionId, messageId);
|
||||||
|
const model = await this.chatSession.get(sessionId).then(s => s?.model);
|
||||||
const provider = this.provider.getProviderByCapability(
|
const provider = this.provider.getProviderByCapability(
|
||||||
(await this.hasAttachment(sessionId, messageId))
|
hasAttachment
|
||||||
? CopilotCapability.ImageToImage
|
? CopilotCapability.ImageToImage
|
||||||
: CopilotCapability.TextToImage
|
: CopilotCapability.TextToImage,
|
||||||
|
model
|
||||||
);
|
);
|
||||||
if (!provider) {
|
if (!provider) {
|
||||||
throw new InternalServerErrorException('No provider available');
|
throw new InternalServerErrorException('No provider available');
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ import {
|
|||||||
CopilotCapability,
|
CopilotCapability,
|
||||||
CopilotImageToImageProvider,
|
CopilotImageToImageProvider,
|
||||||
CopilotProviderType,
|
CopilotProviderType,
|
||||||
|
CopilotTextToImageProvider,
|
||||||
PromptMessage,
|
PromptMessage,
|
||||||
} from '../types';
|
} from '../types';
|
||||||
|
|
||||||
@@ -12,17 +13,24 @@ export type FalConfig = {
|
|||||||
};
|
};
|
||||||
|
|
||||||
export type FalResponse = {
|
export type FalResponse = {
|
||||||
|
detail: Array<{ msg: string }>;
|
||||||
images: Array<{ url: string }>;
|
images: Array<{ url: string }>;
|
||||||
};
|
};
|
||||||
|
|
||||||
export class FalProvider implements CopilotImageToImageProvider {
|
export class FalProvider
|
||||||
|
implements CopilotTextToImageProvider, CopilotImageToImageProvider
|
||||||
|
{
|
||||||
static readonly type = CopilotProviderType.FAL;
|
static readonly type = CopilotProviderType.FAL;
|
||||||
static readonly capabilities = [CopilotCapability.ImageToImage];
|
static readonly capabilities = [
|
||||||
|
CopilotCapability.TextToImage,
|
||||||
|
CopilotCapability.ImageToImage,
|
||||||
|
];
|
||||||
|
|
||||||
readonly availableModels = [
|
readonly availableModels = [
|
||||||
|
// text to image
|
||||||
|
'fast-turbo-diffusion',
|
||||||
// image to image
|
// image to image
|
||||||
// https://blog.fal.ai/building-applications-with-real-time-stable-diffusion-apis/
|
'lcm-sd15-i2i',
|
||||||
'110602490-lcm-sd15-i2i',
|
|
||||||
];
|
];
|
||||||
|
|
||||||
constructor(private readonly config: FalConfig) {
|
constructor(private readonly config: FalConfig) {
|
||||||
@@ -37,6 +45,10 @@ export class FalProvider implements CopilotImageToImageProvider {
|
|||||||
return FalProvider.capabilities;
|
return FalProvider.capabilities;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
isModelAvailable(model: string): boolean {
|
||||||
|
return this.availableModels.includes(model);
|
||||||
|
}
|
||||||
|
|
||||||
// ====== image to image ======
|
// ====== image to image ======
|
||||||
async generateImages(
|
async generateImages(
|
||||||
messages: PromptMessage[],
|
messages: PromptMessage[],
|
||||||
@@ -50,21 +62,20 @@ export class FalProvider implements CopilotImageToImageProvider {
|
|||||||
if (!this.availableModels.includes(model)) {
|
if (!this.availableModels.includes(model)) {
|
||||||
throw new Error(`Invalid model: ${model}`);
|
throw new Error(`Invalid model: ${model}`);
|
||||||
}
|
}
|
||||||
if (!content) {
|
|
||||||
throw new Error('Prompt is required');
|
// prompt attachments require at least one
|
||||||
}
|
if (!content && (!Array.isArray(attachments) || !attachments.length)) {
|
||||||
if (!Array.isArray(attachments) || !attachments.length) {
|
throw new Error('Prompt or Attachments is empty');
|
||||||
throw new Error('Attachments is required');
|
|
||||||
}
|
}
|
||||||
|
|
||||||
const data = (await fetch(`https://${model}.gateway.alpha.fal.ai/`, {
|
const data = (await fetch(`https://fal.run/fal-ai/${model}`, {
|
||||||
method: 'POST',
|
method: 'POST',
|
||||||
headers: {
|
headers: {
|
||||||
Authorization: `key ${this.config.apiKey}`,
|
Authorization: `key ${this.config.apiKey}`,
|
||||||
'Content-Type': 'application/json',
|
'Content-Type': 'application/json',
|
||||||
},
|
},
|
||||||
body: JSON.stringify({
|
body: JSON.stringify({
|
||||||
image_url: attachments[0],
|
image_url: attachments?.[0],
|
||||||
prompt: content,
|
prompt: content,
|
||||||
sync_mode: true,
|
sync_mode: true,
|
||||||
seed: 42,
|
seed: 42,
|
||||||
@@ -73,7 +84,13 @@ export class FalProvider implements CopilotImageToImageProvider {
|
|||||||
signal: options.signal,
|
signal: options.signal,
|
||||||
}).then(res => res.json())) as FalResponse;
|
}).then(res => res.json())) as FalResponse;
|
||||||
|
|
||||||
return data.images.map(image => image.url);
|
if (!data.images?.length) {
|
||||||
|
const error = data.detail?.[0]?.msg;
|
||||||
|
throw new Error(
|
||||||
|
error ? `Invalid message: ${error}` : 'No images generated'
|
||||||
|
);
|
||||||
|
}
|
||||||
|
return data.images?.map(image => image.url) || [];
|
||||||
}
|
}
|
||||||
|
|
||||||
async *generateImagesStream(
|
async *generateImagesStream(
|
||||||
|
|||||||
@@ -118,17 +118,36 @@ export class CopilotProviderService {
|
|||||||
|
|
||||||
getProviderByCapability<C extends CopilotCapability>(
|
getProviderByCapability<C extends CopilotCapability>(
|
||||||
capability: C,
|
capability: C,
|
||||||
|
model?: string,
|
||||||
prefer?: CopilotProviderType
|
prefer?: CopilotProviderType
|
||||||
): CapabilityToCopilotProvider[C] | null {
|
): CapabilityToCopilotProvider[C] | null {
|
||||||
const providers = PROVIDER_CAPABILITY_MAP.get(capability);
|
const providers = PROVIDER_CAPABILITY_MAP.get(capability);
|
||||||
if (Array.isArray(providers) && providers.length) {
|
if (Array.isArray(providers) && providers.length) {
|
||||||
const selectedCapability =
|
let selectedProvider: CopilotProviderType | undefined = prefer;
|
||||||
prefer && providers.includes(prefer) ? prefer : providers[0];
|
let currentIndex = -1;
|
||||||
|
|
||||||
const provider = this.getProvider(selectedCapability);
|
if (!selectedProvider) {
|
||||||
assert(provider.getCapabilities().includes(capability));
|
currentIndex = 0;
|
||||||
|
selectedProvider = providers[currentIndex];
|
||||||
|
}
|
||||||
|
|
||||||
return provider as CapabilityToCopilotProvider[C];
|
while (selectedProvider) {
|
||||||
|
// find first provider that supports the capability and model
|
||||||
|
if (providers.includes(selectedProvider)) {
|
||||||
|
const provider = this.getProvider(selectedProvider);
|
||||||
|
if (provider.getCapabilities().includes(capability)) {
|
||||||
|
if (model) {
|
||||||
|
if (provider.isModelAvailable(model)) {
|
||||||
|
return provider as CapabilityToCopilotProvider[C];
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
return provider as CapabilityToCopilotProvider[C];
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
currentIndex += 1;
|
||||||
|
selectedProvider = providers[currentIndex];
|
||||||
|
}
|
||||||
}
|
}
|
||||||
return null;
|
return null;
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -63,6 +63,10 @@ export class OpenAIProvider
|
|||||||
return OpenAIProvider.capabilities;
|
return OpenAIProvider.capabilities;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
isModelAvailable(model: string): boolean {
|
||||||
|
return this.availableModels.includes(model);
|
||||||
|
}
|
||||||
|
|
||||||
private chatToGPTMessage(
|
private chatToGPTMessage(
|
||||||
messages: PromptMessage[]
|
messages: PromptMessage[]
|
||||||
): OpenAI.Chat.Completions.ChatCompletionMessageParam[] {
|
): OpenAI.Chat.Completions.ChatCompletionMessageParam[] {
|
||||||
|
|||||||
@@ -141,6 +141,7 @@ export enum CopilotCapability {
|
|||||||
|
|
||||||
export interface CopilotProvider {
|
export interface CopilotProvider {
|
||||||
getCapabilities(): CopilotCapability[];
|
getCapabilities(): CopilotCapability[];
|
||||||
|
isModelAvailable(model: string): boolean;
|
||||||
}
|
}
|
||||||
|
|
||||||
export interface CopilotTextToTextProvider extends CopilotProvider {
|
export interface CopilotTextToTextProvider extends CopilotProvider {
|
||||||
|
|||||||
Reference in New Issue
Block a user