feat(server): switch i2i to gpt (#12238)

fix AI-14
fix AI-17
fix AI-39
fix AI-112

<!-- This is an auto-generated comment: release notes by coderabbit.ai -->
## Summary by CodeRabbit

- **New Features**
  - Expanded and reorganized prompt options for text and image actions, adding new prompts for image generation, style conversions, upscaling, background removal, and sticker creation.
  - Enhanced image editing capabilities with direct support for image attachments in prompts.

- **Improvements**
  - Updated prompt names and descriptions to be more user-friendly and descriptive.
  - Simplified and clarified prompt selection and image processing workflows with improved default behaviors.
  - Better organization of prompts through clear grouping and categorization.

- **Bug Fixes**
  - Improved validation and handling of image attachments during editing requests.

- **Refactor**
  - Internal code restructuring of prompts and provider logic for clarity and maintainability without affecting user workflows.
  - Refined message handling and content merging logic to ensure consistent prompt processing.
  - Adjusted image attachment rendering logic for improved display consistency.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->
This commit is contained in:
darkskygit
2025-05-27 11:36:47 +00:00
parent 1e9cbdb65d
commit 3c0fa429c5
10 changed files with 342 additions and 145 deletions

View File

@@ -518,12 +518,7 @@ const actions = [
type: 'text' as const, type: 'text' as const,
}, },
{ {
promptName: [ promptName: ['Convert to sticker', 'Remove background', 'Upscale image'],
'debug:action:fal-face-to-sticker',
'debug:action:fal-remove-bg',
'debug:action:fal-sd15',
'debug:action:fal-upscaler',
],
messages: [ messages: [
{ {
role: 'user' as const, role: 'user' as const,
@@ -590,6 +585,8 @@ for (const {
}))!; }))!;
t.truthy(provider, 'should have provider'); t.truthy(provider, 'should have provider');
await retry(`action: ${promptName}`, t, async t => { await retry(`action: ${promptName}`, t, async t => {
const finalConfig = Object.assign({}, prompt.config, config);
switch (type) { switch (type) {
case 'text': { case 'text': {
const result = await provider.text( const result = await provider.text(
@@ -604,7 +601,7 @@ for (const {
), ),
...messages, ...messages,
], ],
Object.assign({}, prompt.config, config) finalConfig
); );
t.truthy(result, 'should return result'); t.truthy(result, 'should return result');
verifier?.(t, result); verifier?.(t, result);
@@ -622,23 +619,39 @@ for (const {
), ),
...messages, ...messages,
], ],
Object.assign({}, prompt.config, config) finalConfig
); );
t.truthy(result, 'should return result'); t.truthy(result, 'should return result');
verifier?.(t, result); verifier?.(t, result);
break; break;
} }
case 'image': { case 'image': {
const stream = provider.streamImages({ modelId: prompt.model }, [ const finalMessage = [...messages];
const params = {};
if (finalMessage.length === 1) {
const latestMessage = finalMessage.pop()!;
Object.assign(params, {
content: latestMessage.content,
attachments:
'attachments' in latestMessage
? latestMessage.attachments
: undefined,
});
}
const stream = provider.streamImages(
{ modelId: prompt.model },
[
...prompt.finish( ...prompt.finish(
messages.reduce( finalMessage.reduce(
// @ts-expect-error // @ts-expect-error
(acc, m) => Object.assign(acc, m.params), (acc, m) => Object.assign(acc, m.params),
{} params
) )
), ),
...messages, ...finalMessage,
]); ],
finalConfig
);
const result = []; const result = [];
for await (const attachment of stream) { for await (const attachment of stream) {

View File

@@ -543,12 +543,19 @@ test('should be able to chat with special image model', async t => {
); );
}; };
await testWithModel('debug:action:fal-sd15', 'some-tag'); await testWithModel('Generate image', 'some-tag');
await testWithModel( await testWithModel(
'debug:action:fal-upscaler', 'Convert to sticker',
'best quality, 8K resolution, highres, clarity, some-tag' 'convert this image to sticker. you need to identify the subject matter and warp a circle of white stroke around the subject matter and with transparent background. some-tag'
);
await testWithModel(
'Upscale image',
'make the image more detailed. some-tag'
);
await testWithModel(
'Remove background',
'Keep the subject and remove other non-subject items. Transparent background. some-tag'
); );
await testWithModel('debug:action:fal-remove-bg', 'some-tag');
Sinon.restore(); Sinon.restore();
}); });

View File

@@ -84,29 +84,12 @@ export class MockCopilotProvider extends OpenAIProvider {
], ],
}, },
{ {
id: 'lcm-sd15-i2i', id: 'gpt-image-1',
capabilities: [ capabilities: [
{ {
input: [ModelInputType.Image], input: [ModelInputType.Text, ModelInputType.Image],
output: [ModelOutputType.Image],
},
],
},
{
id: 'clarity-upscaler',
capabilities: [
{
input: [ModelInputType.Image],
output: [ModelOutputType.Image],
},
],
},
{
id: 'imageutils/rembg',
capabilities: [
{
input: [ModelInputType.Image],
output: [ModelOutputType.Image], output: [ModelOutputType.Image],
defaultForOutputType: true,
}, },
], ],
}, },

View File

@@ -20,12 +20,6 @@ type Prompt = Omit<
}; };
const workflows: Prompt[] = [ const workflows: Prompt[] = [
{
name: 'debug:action:fal-teed',
action: 'fal-teed',
model: 'workflowutils/teed',
messages: [{ role: 'user', content: '{{content}}' }],
},
{ {
name: 'workflow:presentation', name: 'workflow:presentation',
action: 'workflow:presentation', action: 'workflow:presentation',
@@ -305,48 +299,7 @@ const workflows: Prompt[] = [
}, },
]; ];
const actions: Prompt[] = [ const textActions: Prompt[] = [
{
name: 'debug:action:dalle3',
action: 'image',
model: 'dall-e-3',
messages: [],
},
{
name: 'debug:action:gpt-image-1',
action: 'image',
model: 'gpt-image-1',
messages: [],
},
{
name: 'debug:action:fal-sd15',
action: 'image',
model: 'lcm-sd15-i2i',
messages: [],
},
{
name: 'debug:action:fal-upscaler',
action: 'Clearer',
model: 'clarity-upscaler',
messages: [
{
role: 'user',
content: 'best quality, 8K resolution, highres, clarity, {{content}}',
},
],
},
{
name: 'debug:action:fal-remove-bg',
action: 'Remove background',
model: 'imageutils/rembg',
messages: [],
},
{
name: 'debug:action:fal-face-to-sticker',
action: 'Convert to sticker',
model: 'face-to-sticker',
messages: [],
},
{ {
name: 'Transcript audio', name: 'Transcript audio',
action: 'Transcript audio', action: 'Transcript audio',
@@ -1449,6 +1402,161 @@ When sent new notes, respond ONLY with the contents of the html file.`,
}, },
]; ];
const imageActions: Prompt[] = [
{
name: 'Generate image',
action: 'image',
model: 'gpt-image-1',
messages: [
{
role: 'user',
content: '{{content}}',
},
],
},
{
name: 'Convert to Clay style',
action: 'Convert to Clay style',
model: 'gpt-image-1',
messages: [
{
role: 'user',
content:
'Migration style. Migrates the style from the first image to the second. turn to clay/claymation style. {{content}}',
},
],
},
{
name: 'Convert to Sketch style',
action: 'Convert to Sketch style',
model: 'gpt-image-1',
messages: [
{
role: 'user',
content: 'turn to mono-color sketch style. {{content}}',
},
],
},
{
name: 'Convert to Anime style',
action: 'Convert to Anime style',
model: 'gpt-image-1',
messages: [
{
role: 'user',
content: 'turn to Suzume style like anime style. {{content}}',
},
],
},
{
name: 'Convert to Pixel style',
action: 'Convert to Pixel style',
model: 'gpt-image-1',
messages: [
{
role: 'user',
content: 'turn to kairosoft pixel art. {{content}}',
},
],
},
{
name: 'Convert to sticker',
action: 'Convert to sticker',
model: 'gpt-image-1',
messages: [
{
role: 'user',
content:
'convert this image to sticker. you need to identify the subject matter and warp a circle of white stroke around the subject matter and with transparent background. {{content}}',
},
],
},
{
name: 'Upscale image',
action: 'Upscale image',
model: 'gpt-image-1',
messages: [
{
role: 'user',
content: 'make the image more detailed. {{content}}',
},
],
},
{
name: 'Remove background',
action: 'Remove background',
model: 'gpt-image-1',
messages: [
{
role: 'user',
content:
'Keep the subject and remove other non-subject items. Transparent background. {{content}}',
},
],
},
// TODO(@darkskygit): deprecated, remove it after <0.22 version is outdated
{
name: 'debug:action:fal-remove-bg',
action: 'Remove background',
model: 'imageutils/rembg',
messages: [],
},
{
name: 'debug:action:fal-face-to-sticker',
action: 'Convert to sticker',
model: 'face-to-sticker',
messages: [],
},
{
name: 'debug:action:fal-teed',
action: 'fal-teed',
model: 'workflowutils/teed',
messages: [{ role: 'user', content: '{{content}}' }],
},
{
name: 'debug:action:dalle3',
action: 'image',
model: 'dall-e-3',
messages: [
{
role: 'user',
content: '{{content}}',
},
],
},
{
name: 'debug:action:gpt-image-1',
action: 'image',
model: 'gpt-image-1',
messages: [
{
role: 'user',
content: '{{content}}',
},
],
config: {
requireContent: false,
},
},
{
name: 'debug:action:fal-sd15',
action: 'image',
model: 'lcm-sd15-i2i',
messages: [],
},
{
name: 'debug:action:fal-upscaler',
action: 'Clearer',
model: 'clarity-upscaler',
messages: [
{
role: 'user',
content: 'best quality, 8K resolution, highres, clarity, {{content}}',
},
],
},
];
const CHAT_PROMPT: Omit<Prompt, 'name'> = { const CHAT_PROMPT: Omit<Prompt, 'name'> = {
model: 'gpt-4.1', model: 'gpt-4.1',
optionalModels: [ optionalModels: [
@@ -1622,7 +1730,12 @@ const chat: Prompt[] = [
}, },
]; ];
export const prompts: Prompt[] = [...actions, ...chat, ...workflows]; export const prompts: Prompt[] = [
...textActions,
...imageActions,
...chat,
...workflows,
];
export async function refreshPrompts(db: PrismaClient) { export async function refreshPrompts(db: PrismaClient) {
const needToSkip = await db.aiPrompt const needToSkip = await db.aiPrompt

View File

@@ -13,6 +13,7 @@ import {
streamText, streamText,
ToolSet, ToolSet,
} from 'ai'; } from 'ai';
import { z } from 'zod';
import { import {
CopilotPromptInvalid, CopilotPromptInvalid,
@@ -40,6 +41,20 @@ export type OpenAIConfig = {
baseUrl?: string; baseUrl?: string;
}; };
const ImageResponseSchema = z.union([
z.object({
data: z.array(z.object({ b64_json: z.string() })),
}),
z.object({
error: z.object({
message: z.string(),
type: z.string().nullish(),
param: z.any().nullish(),
code: z.union([z.string(), z.number()]).nullish(),
}),
}),
]);
export class OpenAIProvider extends CopilotProvider<OpenAIConfig> { export class OpenAIProvider extends CopilotProvider<OpenAIConfig> {
readonly type = CopilotProviderType.OpenAI; readonly type = CopilotProviderType.OpenAI;
@@ -389,6 +404,63 @@ export class OpenAIProvider extends CopilotProvider<OpenAIConfig> {
} }
} }
// ====== text to image ======
private async *generateImageWithAttachments(
model: string,
prompt: string,
attachments: NonNullable<PromptMessage['attachments']>
): AsyncGenerator<string> {
const form = new FormData();
form.set('model', model);
form.set('prompt', prompt);
form.set('output_format', 'webp');
for (const [idx, entry] of attachments.entries()) {
const url = typeof entry === 'string' ? entry : entry.attachment;
const resp = await fetch(url);
if (resp.ok) {
const type = resp.headers.get('content-type');
if (type && type.startsWith('image/')) {
const buffer = new Uint8Array(await resp.arrayBuffer());
const file = new File([buffer], `${idx}.png`, { type });
form.append('image[]', file);
}
}
}
if (!form.getAll('image[]').length) {
throw new CopilotPromptInvalid(
'No valid image attachments found. Please attach images.'
);
}
const url = `${this.config.baseUrl || 'https://api.openai.com'}/v1/images/edits`;
const res = await fetch(url, {
method: 'POST',
headers: { Authorization: `Bearer ${this.config.apiKey}` },
body: form,
});
if (!res.ok) {
throw new Error(`OpenAI API error ${res.status}: ${await res.text()}`);
}
const json = await res.json();
const imageResponse = ImageResponseSchema.safeParse(json);
if (imageResponse.success) {
const data = imageResponse.data;
if ('error' in data) {
throw new Error(data.error.message);
} else {
for (const image of data.data) {
yield `data:image/webp;base64,${image.b64_json}`;
}
}
} else {
throw new Error(imageResponse.error.message);
}
}
override async *streamImages( override async *streamImages(
cond: ModelConditions, cond: ModelConditions,
messages: PromptMessage[], messages: PromptMessage[],
@@ -402,12 +474,14 @@ export class OpenAIProvider extends CopilotProvider<OpenAIConfig> {
.counter('generate_images_stream_calls') .counter('generate_images_stream_calls')
.add(1, { model: model.id }); .add(1, { model: model.id });
const { content: prompt } = [...messages].pop() || {}; const { content: prompt, attachments } = [...messages].pop() || {};
if (!prompt) throw new CopilotPromptInvalid('Prompt is required'); if (!prompt) throw new CopilotPromptInvalid('Prompt is required');
try { try {
if (attachments && attachments.length > 0) {
yield* this.generateImageWithAttachments(model.id, prompt, attachments);
} else {
const modelInstance = this.#instance.image(model.id); const modelInstance = this.#instance.image(model.id);
const result = await generateImage({ const result = await generateImage({
model: modelInstance, model: modelInstance,
prompt, prompt,
@@ -428,6 +502,7 @@ export class OpenAIProvider extends CopilotProvider<OpenAIConfig> {
break; break;
} }
} }
}
return; return;
} catch (e: any) { } catch (e: any) {
metrics.ai.counter('generate_images_errors').add(1, { model: model.id }); metrics.ai.counter('generate_images_errors').add(1, { model: model.id });

View File

@@ -39,7 +39,7 @@ const FORMAT_INFER_MAP: Record<string, string> = {
flv: 'video/flv', flv: 'video/flv',
}; };
async function inferMimeType(url: string) { export async function inferMimeType(url: string) {
if (url.startsWith('data:')) { if (url.startsWith('data:')) {
return url.split(';')[0].split(':')[1]; return url.split(';')[0].split(':')[1];
} }

View File

@@ -141,16 +141,13 @@ export class ChatSession implements AsyncDisposable {
return ret; return ret;
} }
finish(params: PromptParams): PromptMessage[] { private mergeUserContent(params: PromptParams) {
const messages = this.takeMessages(); const messages = this.stashMessages;
const firstMessage = messages.at(0); const firstMessage = messages.at(0);
// TODO: refactor this {{content}} keyword agreement
// if the message in prompt config contains {{content}},
// we should combine it with the user message in the prompt
if ( if (
messages.length === 1 && this.state.prompt.paramKeys.includes('content') &&
firstMessage && !messages.some(m => m.role === AiPromptRole.assistant) &&
this.state.prompt.paramKeys.includes('content') firstMessage
) { ) {
const normalizedParams = { const normalizedParams = {
...params, ...params,
@@ -178,7 +175,18 @@ export class ChatSession implements AsyncDisposable {
return finished; return finished;
} }
return;
}
finish(params: PromptParams): PromptMessage[] {
// if the message in prompt config contains {{content}},
// we should combine it with the user message in the prompt
const mergedMessage = this.mergeUserContent(params);
if (mergedMessage) {
return mergedMessage;
}
const messages = this.takeMessages();
const lastMessage = messages.at(-1); const lastMessage = messages.at(-1);
return [ return [
...this.state.prompt.finish( ...this.state.prompt.finish(

View File

@@ -21,7 +21,8 @@ export class ActionImage extends WithDisposable(ShadowlessElement) {
accessor testId = 'action-image'; accessor testId = 'action-image';
protected override render() { protected override render() {
const images = this.item.messages[0].attachments; const images =
this.item.messages[1]?.attachments ?? this.item.messages[0].attachments;
return html`<action-wrapper .host=${this.host} .item=${this.item}> return html`<action-wrapper .host=${this.host} .item=${this.item}>
<div style=${styleMap({ marginBottom: '12px' })}> <div style=${styleMap({ marginBottom: '12px' })}>

View File

@@ -1,12 +1,7 @@
// manually synced with packages/backend/server/src/data/migrations/utils/prompts.ts // manually synced with packages/backend/server/src/data/migrations/utils/prompts.ts
// TODO(@Peng): automate this // TODO(@Peng): automate this
export const promptKeys = [ export const promptKeys = [
'debug:action:dalle3', // text actions
'debug:action:gpt-image-1',
'debug:action:fal-sd15',
'debug:action:fal-upscaler',
'debug:action:fal-remove-bg',
'debug:action:fal-face-to-sticker',
'Chat With AFFiNE AI', 'Chat With AFFiNE AI',
'Search With AFFiNE AI', 'Search With AFFiNE AI',
'Summary', 'Summary',
@@ -36,12 +31,18 @@ export const promptKeys = [
'Make it longer', 'Make it longer',
'Make it shorter', 'Make it shorter',
'Continue writing', 'Continue writing',
// image actions
'Generate image',
'Convert to Anime style',
'Convert to Clay style',
'Convert to Pixel style',
'Convert to Sketch style',
'Convert to sticker',
'Upscale image',
'Remove background',
// workflows
'workflow:presentation', 'workflow:presentation',
'workflow:brainstorm', 'workflow:brainstorm',
'workflow:image-sketch',
'workflow:image-clay',
'workflow:image-anime',
'workflow:image-pixel',
] as const; ] as const;
export type PromptKey = (typeof promptKeys)[number]; export type PromptKey = (typeof promptKeys)[number];

View File

@@ -27,18 +27,18 @@ function toAIUserInfo(account: AuthAccountInfo | null) {
const filterStyleToPromptName = new Map<string, PromptKey>( const filterStyleToPromptName = new Map<string, PromptKey>(
Object.entries({ Object.entries({
'Clay style': 'workflow:image-clay', 'Clay style': 'Convert to Clay style',
'Pixel style': 'workflow:image-pixel', 'Pixel style': 'Convert to Pixel style',
'Sketch style': 'workflow:image-sketch', 'Sketch style': 'Convert to Sketch style',
'Anime style': 'workflow:image-anime', 'Anime style': 'Convert to Anime style',
}) })
); );
const processTypeToPromptName = new Map<string, PromptKey>( const processTypeToPromptName = new Map<string, PromptKey>(
Object.entries({ Object.entries({
Clearer: 'debug:action:fal-upscaler', Clearer: 'Upscale image',
'Remove background': 'debug:action:fal-remove-bg', 'Remove background': 'Remove background',
'Convert to sticker': 'debug:action:fal-face-to-sticker', 'Convert to sticker': 'Convert to sticker',
}) })
); );
@@ -486,22 +486,18 @@ Could you make a new website based on these notes and send back just the html fi
}); });
AIProvider.provide('createImage', async options => { AIProvider.provide('createImage', async options => {
// test to image
let promptName: PromptKey = 'debug:action:gpt-image-1';
// image to image
if (options.attachments?.length) {
promptName = 'debug:action:fal-sd15';
}
const sessionId = await createSession({ const sessionId = await createSession({
promptName, promptName: 'Generate image',
...options, ...options,
}); });
return toImage({ return toImage({
...options, ...options,
client, client,
sessionId, sessionId,
content: options.input, content:
!options.input && options.attachments
? 'Make the image more detailed.'
: options.input,
// 5 minutes // 5 minutes
timeout: 300000, timeout: 300000,
}); });