fix: pick copilot provider depend on model (#6540)

This commit is contained in:
darkskygit
2024-04-12 12:01:39 +00:00
parent 62f90e5f10
commit fc51b68674
6 changed files with 75 additions and 21 deletions

View File

@@ -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',

View File

@@ -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');

View File

@@ -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(

View File

@@ -118,18 +118,37 @@ 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];
}
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]; return provider as CapabilityToCopilotProvider[C];
} }
} else {
return provider as CapabilityToCopilotProvider[C];
}
}
}
currentIndex += 1;
selectedProvider = providers[currentIndex];
}
}
return null; return null;
} }
} }

View File

@@ -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[] {

View File

@@ -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 {