refactor(core): ai session create (#11539)

Close [BS-3079](https://linear.app/affine-design/issue/BS-3079).

- Separate the create session logic from the `createMessage function`.
- Ensure the session is created before executing any chat or actions.
- Convert the `AIActions` into asynchronous functions.
- Transfer the update prompt name logic to the chat action.
- Introduce a networkSearch field in `AITextActionOptions`.
- Eliminate the redundant `LAST_ROOT_SESSION_ID`.
This commit is contained in:
akumatus
2025-04-09 06:18:56 +00:00
parent 1e9561b46c
commit c9790ed854
14 changed files with 403 additions and 331 deletions

View File

@@ -743,7 +743,7 @@ export class MindmapElementModel extends GfxGroupLikeElementModel<MindmapElement
const targetPos = const targetPos =
typeof targetXYWH === 'string' ? deserializeXYWH(targetXYWH) : targetXYWH; typeof targetXYWH === 'string' ? deserializeXYWH(targetXYWH) : targetXYWH;
const offsetX = targetPos[0] - x; const offsetX = targetPos[0] - x;
const offsetY = targetPos[1] - y + targetPos[3]; const offsetY = targetPos[1] - y;
this.surface.doc.transact(() => { this.surface.doc.transact(() => {
this.childElements.forEach(el => { this.childElements.forEach(el => {

View File

@@ -108,7 +108,7 @@ function actionToStream<T extends keyof BlockSuitePresets.AIActions>(
workspaceId: host.doc.workspace.id, workspaceId: host.doc.workspace.id,
} as Parameters<typeof action>[0]; } as Parameters<typeof action>[0];
// @ts-expect-error TODO(@Peng): maybe fix this // @ts-expect-error TODO(@Peng): maybe fix this
stream = action(options); stream = await action(options);
if (!stream) return; if (!stream) return;
yield* stream; yield* stream;
}, },

View File

@@ -207,7 +207,7 @@ function actionToStream<T extends keyof BlockSuitePresets.AIActions>(
} }
// @ts-expect-error TODO(@Peng): maybe fix this // @ts-expect-error TODO(@Peng): maybe fix this
stream = action(options); stream = await action(options);
if (!stream) return; if (!stream) return;
yield* stream; yield* stream;
}, },
@@ -237,7 +237,7 @@ function actionToStream<T extends keyof BlockSuitePresets.AIActions>(
} as Parameters<typeof action>[0]; } as Parameters<typeof action>[0];
// @ts-expect-error TODO(@Peng): maybe fix this // @ts-expect-error TODO(@Peng): maybe fix this
stream = action(options); stream = await action(options);
if (!stream) return; if (!stream) return;
yield* stream; yield* stream;
}, },

View File

@@ -5,7 +5,11 @@ import {
} from '@blocksuite/affine/blocks/surface'; } from '@blocksuite/affine/blocks/surface';
import { fitContent } from '@blocksuite/affine/gfx/shape'; import { fitContent } from '@blocksuite/affine/gfx/shape';
import { createTemplateJob } from '@blocksuite/affine/gfx/template'; import { createTemplateJob } from '@blocksuite/affine/gfx/template';
import { Bound, getCommonBound } from '@blocksuite/affine/global/gfx'; import {
Bound,
getCommonBound,
type XYWH,
} from '@blocksuite/affine/global/gfx';
import type { import type {
MindmapElementModel, MindmapElementModel,
ShapeElementModel, ShapeElementModel,
@@ -83,7 +87,10 @@ function responseToBrainstormMindmap(
}); });
// wait for mindmap xywh update // wait for mindmap xywh update
setTimeout(() => { setTimeout(() => {
const frameBound = expandBound(mindmap.elementBound, PADDING); const { x, y, w, h } = mindmap.elementBound;
const targetBound: XYWH = [x, y + h / 2 + PADDING - 15, w, h];
mindmap.moveTo(targetBound);
const frameBound = expandBound(new Bound(...targetBound), PADDING);
addSurfaceRefBlock(host, frameBound, place); addSurfaceRefBlock(host, frameBound, place);
}, 0); }, 0);
}); });
@@ -107,7 +114,7 @@ function responseToMakeItReal(host: EditorHost, ctx: AIContext, place: Place) {
const bound = getEdgelessContentBound(host); const bound = getEdgelessContentBound(host);
const x = bound ? bound.x + bound.w + PADDING * 2 : 0; const x = bound ? bound.x + bound.w + PADDING * 2 : 0;
const y = bound ? bound.y : 0; const y = bound ? bound.y : 0;
const htmlBound = new Bound(x, y, width || 800, height || 600); const htmlBound = new Bound(x, y + PADDING, width || 800, height || 600);
const html = preprocessHtml(aiPanel.answer); const html = preprocessHtml(aiPanel.answer);
host.doc.transact(() => { host.doc.transact(() => {
host.doc.addBlock( host.doc.addBlock(

View File

@@ -13,6 +13,8 @@ import type { EditorHost } from '@blocksuite/affine/std';
import type { GfxModel } from '@blocksuite/affine/std/gfx'; import type { GfxModel } from '@blocksuite/affine/std/gfx';
import type { BlockModel } from '@blocksuite/affine/store'; import type { BlockModel } from '@blocksuite/affine/store';
import type { PromptKey } from '../provider/prompt';
export const translateLangs = [ export const translateLangs = [
'English', 'English',
'Spanish', 'Spanish',
@@ -131,6 +133,7 @@ declare global {
interface ChatOptions extends AITextActionOptions { interface ChatOptions extends AITextActionOptions {
sessionId?: string; sessionId?: string;
isRootSession?: boolean; isRootSession?: boolean;
networkSearch?: boolean;
contexts?: { contexts?: {
docs: AIDocContextOption[]; docs: AIDocContextOption[];
files: AIFileContextOption[]; files: AIFileContextOption[];
@@ -155,107 +158,107 @@ declare global {
interface AIActions { interface AIActions {
// chat is a bit special because it's has a internally maintained session // chat is a bit special because it's has a internally maintained session
chat<T extends ChatOptions>(options: T): AIActionTextResponse<T>; chat<T extends ChatOptions>(options: T): Promise<AIActionTextResponse<T>>;
summary<T extends AITextActionOptions>( summary<T extends AITextActionOptions>(
options: T options: T
): AIActionTextResponse<T>; ): Promise<AIActionTextResponse<T>>;
improveWriting<T extends AITextActionOptions>( improveWriting<T extends AITextActionOptions>(
options: T options: T
): AIActionTextResponse<T>; ): Promise<AIActionTextResponse<T>>;
improveGrammar<T extends AITextActionOptions>( improveGrammar<T extends AITextActionOptions>(
options: T options: T
): AIActionTextResponse<T>; ): Promise<AIActionTextResponse<T>>;
fixSpelling<T extends AITextActionOptions>( fixSpelling<T extends AITextActionOptions>(
options: T options: T
): AIActionTextResponse<T>; ): Promise<AIActionTextResponse<T>>;
createHeadings<T extends AITextActionOptions>( createHeadings<T extends AITextActionOptions>(
options: T options: T
): AIActionTextResponse<T>; ): Promise<AIActionTextResponse<T>>;
makeLonger<T extends AITextActionOptions>( makeLonger<T extends AITextActionOptions>(
options: T options: T
): AIActionTextResponse<T>; ): Promise<AIActionTextResponse<T>>;
makeShorter<T extends AITextActionOptions>( makeShorter<T extends AITextActionOptions>(
options: T options: T
): AIActionTextResponse<T>; ): Promise<AIActionTextResponse<T>>;
continueWriting<T extends AITextActionOptions>( continueWriting<T extends AITextActionOptions>(
options: T options: T
): AIActionTextResponse<T>; ): Promise<AIActionTextResponse<T>>;
checkCodeErrors<T extends AITextActionOptions>( checkCodeErrors<T extends AITextActionOptions>(
options: T options: T
): AIActionTextResponse<T>; ): Promise<AIActionTextResponse<T>>;
explainCode<T extends AITextActionOptions>( explainCode<T extends AITextActionOptions>(
options: T options: T
): AIActionTextResponse<T>; ): Promise<AIActionTextResponse<T>>;
writeArticle<T extends AITextActionOptions>( writeArticle<T extends AITextActionOptions>(
options: T options: T
): AIActionTextResponse<T>; ): Promise<AIActionTextResponse<T>>;
writeTwitterPost<T extends AITextActionOptions>( writeTwitterPost<T extends AITextActionOptions>(
options: T options: T
): AIActionTextResponse<T>; ): Promise<AIActionTextResponse<T>>;
writePoem<T extends AITextActionOptions>( writePoem<T extends AITextActionOptions>(
options: T options: T
): AIActionTextResponse<T>; ): Promise<AIActionTextResponse<T>>;
writeBlogPost<T extends AITextActionOptions>( writeBlogPost<T extends AITextActionOptions>(
options: T options: T
): AIActionTextResponse<T>; ): Promise<AIActionTextResponse<T>>;
brainstorm<T extends AITextActionOptions>( brainstorm<T extends AITextActionOptions>(
options: T options: T
): AIActionTextResponse<T>; ): Promise<AIActionTextResponse<T>>;
writeOutline<T extends AITextActionOptions>( writeOutline<T extends AITextActionOptions>(
options: T options: T
): AIActionTextResponse<T>; ): Promise<AIActionTextResponse<T>>;
explainImage<T extends AITextActionOptions>( explainImage<T extends AITextActionOptions>(
options: T options: T
): AIActionTextResponse<T>; ): Promise<AIActionTextResponse<T>>;
findActions<T extends AITextActionOptions>( findActions<T extends AITextActionOptions>(
options: T options: T
): AIActionTextResponse<T>; ): Promise<AIActionTextResponse<T>>;
// mindmap // mindmap
brainstormMindmap<T extends BrainstormMindMap>( brainstormMindmap<T extends BrainstormMindMap>(
options: T options: T
): AIActionTextResponse<T>; ): Promise<AIActionTextResponse<T>>;
expandMindmap<T extends ExpandMindMap>( expandMindmap<T extends ExpandMindMap>(
options: T options: T
): AIActionTextResponse<T>; ): Promise<AIActionTextResponse<T>>;
// presentation // presentation
createSlides<T extends AITextActionOptions>( createSlides<T extends AITextActionOptions>(
options: T options: T
): AIActionTextResponse<T>; ): Promise<AIActionTextResponse<T>>;
// explain this // explain this
explain<T extends AITextActionOptions>( explain<T extends AITextActionOptions>(
options: T options: T
): AIActionTextResponse<T>; ): Promise<AIActionTextResponse<T>>;
// actions with variants // actions with variants
translate<T extends TranslateOptions>( translate<T extends TranslateOptions>(
options: T options: T
): AIActionTextResponse<T>; ): Promise<AIActionTextResponse<T>>;
changeTone<T extends ChangeToneOptions>( changeTone<T extends ChangeToneOptions>(
options: T options: T
): AIActionTextResponse<T>; ): Promise<AIActionTextResponse<T>>;
// make it real, image to text // make it real, image to text
makeItReal<T extends AIImageActionOptions>( makeItReal<T extends AIImageActionOptions>(
options: T options: T
): AIActionTextResponse<T>; ): Promise<AIActionTextResponse<T>>;
createImage<T extends AIImageActionOptions>( createImage<T extends AIImageActionOptions>(
options: T options: T
): AIActionTextResponse<T>; ): Promise<AIActionTextResponse<T>>;
processImage<T extends ProcessImageOptions>( processImage<T extends ProcessImageOptions>(
options: T options: T
): AIActionTextResponse<T>; ): Promise<AIActionTextResponse<T>>;
filterImage<T extends FilterImageOptions>( filterImage<T extends FilterImageOptions>(
options: T options: T
): AIActionTextResponse<T>; ): Promise<AIActionTextResponse<T>>;
generateCaption<T extends AITextActionOptions>( generateCaption<T extends AITextActionOptions>(
options: T options: T
): AIActionTextResponse<T>; ): Promise<AIActionTextResponse<T>>;
} }
type AIDocsAndFilesContext = { type AIDocsAndFilesContext = {
@@ -357,12 +360,16 @@ declare global {
>[]; >[];
}; };
interface CreateSessionOptions {
docId: string;
workspaceId: string;
promptName: PromptKey;
sessionId?: string;
retry?: boolean;
}
interface AISessionService { interface AISessionService {
createSession: ( createSession: (options: CreateSessionOptions) => Promise<string>;
workspaceId: string,
docId: string,
promptName?: string
) => Promise<string>;
getSessions: ( getSessions: (
workspaceId: string, workspaceId: string,
docId?: string, docId?: string,

View File

@@ -348,10 +348,10 @@ export class ChatPanelMessages extends WithDisposable(ShadowlessElement) {
} }
retry = async () => { retry = async () => {
const { doc } = this.host;
try { try {
const sessionId = await this.createSessionId(); const sessionId = await this.createSessionId();
if (!sessionId) return; if (!sessionId) return;
if (!AIProvider.actions.chat) return;
const abortController = new AbortController(); const abortController = new AbortController();
const messages = [...this.chatContextValue.messages]; const messages = [...this.chatContextValue.messages];
@@ -362,7 +362,8 @@ export class ChatPanelMessages extends WithDisposable(ShadowlessElement) {
} }
this.updateContext({ messages, status: 'loading', error: null }); this.updateContext({ messages, status: 'loading', error: null });
const stream = AIProvider.actions.chat?.({ const { doc } = this.host;
const stream = await AIProvider.actions.chat({
sessionId, sessionId,
retry: true, retry: true,
docId: doc.id, docId: doc.id,
@@ -374,8 +375,6 @@ export class ChatPanelMessages extends WithDisposable(ShadowlessElement) {
control: 'chat-send', control: 'chat-send',
isRootSession: true, isRootSession: true,
}); });
if (stream) {
this.updateContext({ abortController }); this.updateContext({ abortController });
for await (const text of stream) { for await (const text of stream) {
const messages = [...this.chatContextValue.messages]; const messages = [...this.chatContextValue.messages];
@@ -385,7 +384,6 @@ export class ChatPanelMessages extends WithDisposable(ShadowlessElement) {
} }
this.updateContext({ status: 'success' }); this.updateContext({ status: 'success' });
}
} catch (error) { } catch (error) {
this.updateContext({ status: 'error', error: error as AIError }); this.updateContext({ status: 'error', error: error as AIError });
} finally { } finally {

View File

@@ -141,7 +141,6 @@ export class ChatPanel extends SignalWatcher(
const history = histories?.find(history => history.sessionId === sessionId); const history = histories?.find(history => history.sessionId === sessionId);
if (history) { if (history) {
messages.push(...history.messages); messages.push(...history.messages);
AIProvider.LAST_ROOT_SESSION_ID = history.sessionId;
} }
this.chatContextValue = { this.chatContextValue = {
@@ -182,10 +181,11 @@ export class ChatPanel extends SignalWatcher(
if (this._sessionId) { if (this._sessionId) {
return this._sessionId; return this._sessionId;
} }
this._sessionId = await AIProvider.session?.createSession( this._sessionId = await AIProvider.session?.createSession({
this.doc.workspace.id, docId: this.doc.id,
this.doc.id workspaceId: this.doc.workspace.id,
); promptName: 'Chat With AFFiNE AI',
});
return this._sessionId; return this._sessionId;
}; };

View File

@@ -25,7 +25,6 @@ import type {
} from '../ai-chat-chips/type'; } from '../ai-chat-chips/type';
import { isDocChip, isFileChip } from '../ai-chat-chips/utils'; import { isDocChip, isFileChip } from '../ai-chat-chips/utils';
import type { ChatMessage } from '../ai-chat-messages'; import type { ChatMessage } from '../ai-chat-messages';
import { PROMPT_NAME_AFFINE_AI, PROMPT_NAME_NETWORK_SEARCH } from './const';
import type { AIChatInputContext, AINetworkSearchConfig } from './type'; import type { AIChatInputContext, AINetworkSearchConfig } from './type';
const MaximumImageCount = 32; const MaximumImageCount = 32;
@@ -255,7 +254,8 @@ export class AIChatInput extends SignalWatcher(WithDisposable(LitElement)) {
private get _isNetworkActive() { private get _isNetworkActive() {
return ( return (
!!this.networkSearchConfig.visible.value && !!this.networkSearchConfig.visible.value &&
!!this.networkSearchConfig.enabled.value !!this.networkSearchConfig.enabled.value &&
!this._isNetworkDisabled
); );
} }
@@ -274,22 +274,6 @@ export class AIChatInput extends SignalWatcher(WithDisposable(LitElement)) {
); );
} }
private _getPromptName() {
if (this._isNetworkDisabled) {
return PROMPT_NAME_AFFINE_AI;
}
return this._isNetworkActive
? PROMPT_NAME_NETWORK_SEARCH
: PROMPT_NAME_AFFINE_AI;
}
private async _updatePromptName(promptName: string) {
const sessionId = await this.createSessionId();
if (sessionId && AIProvider.session) {
await AIProvider.session.updateSession(sessionId, promptName);
}
}
override connectedCallback() { override connectedCallback() {
super.connectedCallback(); super.connectedCallback();
this._disposables.add( this._disposables.add(
@@ -313,7 +297,7 @@ export class AIChatInput extends SignalWatcher(WithDisposable(LitElement)) {
const { images, status } = this.chatContextValue; const { images, status } = this.chatContextValue;
const hasImages = images.length > 0; const hasImages = images.length > 0;
const maxHeight = hasImages ? 272 + 2 : 200 + 2; const maxHeight = hasImages ? 272 + 2 : 200 + 2;
const uploadDisabled = this._isNetworkActive && !this._isNetworkDisabled; const uploadDisabled = this._isNetworkActive;
return html` <div return html` <div
class="chat-panel-input" class="chat-panel-input"
data-if-focused=${this.focused} data-if-focused=${this.focused}
@@ -380,9 +364,7 @@ export class AIChatInput extends SignalWatcher(WithDisposable(LitElement)) {
data-testid="chat-network-search" data-testid="chat-network-search"
aria-disabled=${this._isNetworkDisabled} aria-disabled=${this._isNetworkDisabled}
data-active=${this._isNetworkActive} data-active=${this._isNetworkActive}
@click=${this._isNetworkDisabled @click=${this._toggleNetworkSearch}
? undefined
: this._toggleNetworkSearch}
@pointerdown=${stopPropagation} @pointerdown=${stopPropagation}
> >
${PublishIcon()} ${PublishIcon()}
@@ -473,6 +455,9 @@ export class AIChatInput extends SignalWatcher(WithDisposable(LitElement)) {
e.preventDefault(); e.preventDefault();
e.stopPropagation(); e.stopPropagation();
if (this._isNetworkDisabled) {
return;
}
const enable = this.networkSearchConfig.enabled.value; const enable = this.networkSearchConfig.enabled.value;
this.networkSearchConfig.setEnabled(!enable); this.networkSearchConfig.setEnabled(!enable);
}; };
@@ -514,13 +499,12 @@ export class AIChatInput extends SignalWatcher(WithDisposable(LitElement)) {
}; };
send = async (text: string) => { send = async (text: string) => {
try {
const { status, markdown, images } = this.chatContextValue; const { status, markdown, images } = this.chatContextValue;
if (status === 'loading' || status === 'transmitting') return; if (status === 'loading' || status === 'transmitting') return;
if (!text) return; if (!text) return;
if (!AIProvider.actions.chat) return; if (!AIProvider.actions.chat) return;
try {
const promptName = this._getPromptName();
const abortController = new AbortController(); const abortController = new AbortController();
this.updateContext({ this.updateContext({
images: [], images: [],
@@ -538,16 +522,13 @@ export class AIChatInput extends SignalWatcher(WithDisposable(LitElement)) {
// optimistic update messages // optimistic update messages
await this._preUpdateMessages(userInput, attachments); await this._preUpdateMessages(userInput, attachments);
// must update prompt name after local chat message is updated
// otherwise, the unauthorized error can not be rendered properly
await this._updatePromptName(promptName);
const sessionId = await this.createSessionId(); const sessionId = await this.createSessionId();
const contexts = await this._getMatchedContexts(userInput); const contexts = await this._getMatchedContexts(userInput);
if (abortController.signal.aborted) { if (abortController.signal.aborted) {
return; return;
} }
const stream = AIProvider.actions.chat({ const stream = await AIProvider.actions.chat({
sessionId, sessionId,
input: userInput, input: userInput,
contexts, contexts,
@@ -560,6 +541,7 @@ export class AIChatInput extends SignalWatcher(WithDisposable(LitElement)) {
isRootSession: this.isRootSession, isRootSession: this.isRootSession,
where: this.trackOptions.where, where: this.trackOptions.where,
control: this.trackOptions.control, control: this.trackOptions.control,
networkSearch: this._isNetworkActive,
}); });
for await (const text of stream) { for await (const text of stream) {

View File

@@ -1,2 +0,0 @@
export const PROMPT_NAME_AFFINE_AI = 'Chat With AFFiNE AI';
export const PROMPT_NAME_NETWORK_SEARCH = 'Search With AFFiNE AI';

View File

@@ -1,3 +1,2 @@
export * from './ai-chat-input'; export * from './ai-chat-input';
export * from './const';
export * from './type'; export * from './type';

View File

@@ -320,16 +320,12 @@ export class AIChatBlockPeekView extends LitElement {
* Retry the last chat message * Retry the last chat message
*/ */
retry = async () => { retry = async () => {
const { doc } = this.host;
const { _forkBlockId, _forkSessionId } = this;
if (!_forkBlockId || !_forkSessionId) {
return;
}
let content = '';
try { try {
const abortController = new AbortController(); const { _forkBlockId, _forkSessionId } = this;
if (!_forkBlockId || !_forkSessionId) return;
if (!AIProvider.actions.chat) return;
const abortController = new AbortController();
const messages = [...this.chatContext.messages]; const messages = [...this.chatContext.messages];
const last = messages[messages.length - 1]; const last = messages[messages.length - 1];
if ('content' in last) { if ('content' in last) {
@@ -339,7 +335,8 @@ export class AIChatBlockPeekView extends LitElement {
} }
this.updateContext({ messages, status: 'loading', error: null }); this.updateContext({ messages, status: 'loading', error: null });
const stream = AIProvider.actions.chat?.({ const { doc } = this.host;
const stream = await AIProvider.actions.chat({
sessionId: _forkSessionId, sessionId: _forkSessionId,
retry: true, retry: true,
docId: doc.id, docId: doc.id,
@@ -351,26 +348,21 @@ export class AIChatBlockPeekView extends LitElement {
control: 'chat-send', control: 'chat-send',
}); });
if (stream) {
this.updateContext({ abortController }); this.updateContext({ abortController });
for await (const text of stream) { for await (const text of stream) {
const messages = [...this.chatContext.messages]; const messages = [...this.chatContext.messages];
const last = messages[messages.length - 1] as ChatMessage; const last = messages[messages.length - 1] as ChatMessage;
last.content += text; last.content += text;
this.updateContext({ messages, status: 'transmitting' }); this.updateContext({ messages, status: 'transmitting' });
content += text;
} }
this.updateContext({ status: 'success' }); this.updateContext({ status: 'success' });
} // Update new chat block messages if there are contents returned from AI
await this.updateChatBlockMessages();
} catch (error) { } catch (error) {
this.updateContext({ status: 'error', error: error as AIError }); this.updateContext({ status: 'error', error: error as AIError });
} finally { } finally {
this.updateContext({ abortController: null }); this.updateContext({ abortController: null });
if (content) {
// Update new chat block messages if there are contents returned from AI
await this.updateChatBlockMessages();
}
} }
}; };

View File

@@ -100,8 +100,6 @@ export class AIProvider {
static LAST_ACTION_SESSIONID = ''; static LAST_ACTION_SESSIONID = '';
static LAST_ROOT_SESSION_ID = '';
static MAX_LOCAL_HISTORY = 10; static MAX_LOCAL_HISTORY = 10;
private readonly actions: Partial<BlockSuitePresets.AIActions> = {}; private readonly actions: Partial<BlockSuitePresets.AIActions> = {};
@@ -158,10 +156,10 @@ export class AIProvider {
id: T, id: T,
action: ( action: (
...options: Parameters<BlockSuitePresets.AIActions[T]> ...options: Parameters<BlockSuitePresets.AIActions[T]>
) => ReturnType<BlockSuitePresets.AIActions[T]> ) => Promise<ReturnType<BlockSuitePresets.AIActions[T]>>
): void { ): void {
// @ts-expect-error TODO: maybe fix this // @ts-expect-error TODO: maybe fix this
this.actions[id] = ( this.actions[id] = async (
...args: Parameters<BlockSuitePresets.AIActions[T]> ...args: Parameters<BlockSuitePresets.AIActions[T]>
) => { ) => {
const options = args[0]; const options = args[0];
@@ -176,9 +174,8 @@ export class AIProvider {
this.actionHistory.shift(); this.actionHistory.shift();
} }
// wrap the action with slot actions // wrap the action with slot actions
const result: BlockSuitePresets.TextStream | Promise<string> = action( const result: BlockSuitePresets.TextStream | Promise<string> =
...args await action(...args);
);
const isTextStream = ( const isTextStream = (
m: BlockSuitePresets.TextStream | Promise<string> m: BlockSuitePresets.TextStream | Promise<string>
): m is BlockSuitePresets.TextStream => ): m is BlockSuitePresets.TextStream =>
@@ -315,7 +312,7 @@ export class AIProvider {
id: T, id: T,
action: ( action: (
...options: Parameters<BlockSuitePresets.AIActions[T]> ...options: Parameters<BlockSuitePresets.AIActions[T]>
) => ReturnType<BlockSuitePresets.AIActions[T]> ) => Promise<ReturnType<BlockSuitePresets.AIActions[T]>>
): void; ): void;
static provide(id: unknown, action: unknown) { static provide(id: unknown, action: unknown) {

View File

@@ -3,16 +3,12 @@ import { partition } from 'lodash-es';
import { AIProvider } from './ai-provider'; import { AIProvider } from './ai-provider';
import type { CopilotClient } from './copilot-client'; import type { CopilotClient } from './copilot-client';
import { delay, toTextStream } from './event-source'; import { delay, toTextStream } from './event-source';
import type { PromptKey } from './prompt';
const TIMEOUT = 50000; const TIMEOUT = 50000;
export type TextToTextOptions = { export type TextToTextOptions = {
client: CopilotClient; client: CopilotClient;
docId: string; sessionId: string;
workspaceId: string;
promptName?: PromptKey;
sessionId?: string | Promise<string>;
content?: string; content?: string;
attachments?: (string | Blob | File)[]; attachments?: (string | Blob | File)[];
params?: Record<string, any>; params?: Record<string, any>;
@@ -61,30 +57,22 @@ async function resizeImage(blob: Blob | File): Promise<Blob | null> {
return null; return null;
} }
async function createSessionMessage({ interface CreateMessageOptions {
client: CopilotClient;
sessionId: string;
content?: string;
attachments?: (string | Blob | File)[];
params?: Record<string, any>;
}
async function createMessage({
client, client,
docId, sessionId,
workspaceId,
promptName = 'Chat With AFFiNE AI',
content, content,
sessionId: providedSessionId,
attachments, attachments,
params, params,
}: TextToTextOptions): Promise<{ }: CreateMessageOptions): Promise<string> {
sessionId: string;
messageId: string;
}> {
if (!promptName && !providedSessionId) {
throw new Error('promptName or sessionId is required');
}
const hasAttachments = attachments && attachments.length > 0; const hasAttachments = attachments && attachments.length > 0;
const sessionId = await (providedSessionId ??
client.createSession({
workspaceId,
docId,
promptName,
}));
const options: Parameters<CopilotClient['createMessage']>[0] = { const options: Parameters<CopilotClient['createMessage']>[0] = {
sessionId, sessionId,
content, content,
@@ -110,67 +98,44 @@ async function createSessionMessage({
).filter(Boolean) as File[]; ).filter(Boolean) as File[];
} }
const messageId = await client.createMessage(options); return await client.createMessage(options);
return {
messageId,
sessionId,
};
} }
export function textToText({ export function textToText({
client, client,
docId, sessionId,
workspaceId,
promptName,
content, content,
attachments, attachments,
params, params,
sessionId,
stream, stream,
signal, signal,
timeout = TIMEOUT, timeout = TIMEOUT,
retry = false, retry = false,
workflow = false, workflow = false,
isRootSession = false,
postfix, postfix,
}: TextToTextOptions) { }: TextToTextOptions) {
let _sessionId: string; let messageId: string | undefined;
let _messageId: string | undefined;
if (stream) { if (stream) {
return { return {
[Symbol.asyncIterator]: async function* () { [Symbol.asyncIterator]: async function* () {
if (retry) { if (!retry) {
const retrySessionId = messageId = await createMessage({
(await sessionId) ?? AIProvider.LAST_ACTION_SESSIONID;
_sessionId = retrySessionId;
_messageId = undefined;
} else {
const message = await createSessionMessage({
client, client,
docId, sessionId,
workspaceId,
promptName,
content, content,
attachments, attachments,
params, params,
sessionId,
}); });
_sessionId = message.sessionId;
_messageId = message.messageId;
} }
const eventSource = client.chatTextStream( const eventSource = client.chatTextStream(
{ {
sessionId: _sessionId, sessionId,
messageId: _messageId, messageId,
}, },
workflow ? 'workflow' : undefined workflow ? 'workflow' : undefined
); );
AIProvider.LAST_ACTION_SESSIONID = _sessionId; AIProvider.LAST_ACTION_SESSIONID = sessionId;
if (isRootSession) {
AIProvider.LAST_ROOT_SESSION_ID = _sessionId;
}
if (signal) { if (signal) {
if (signal.aborted) { if (signal.aborted) {
@@ -212,34 +177,20 @@ export function textToText({
}) })
: null, : null,
(async function () { (async function () {
if (retry) { if (!retry) {
const retrySessionId = messageId = await createMessage({
(await sessionId) ?? AIProvider.LAST_ACTION_SESSIONID;
_sessionId = retrySessionId;
_messageId = undefined;
} else {
const message = await createSessionMessage({
client, client,
docId, sessionId,
workspaceId,
promptName,
content, content,
attachments, attachments,
params, params,
sessionId,
}); });
_sessionId = message.sessionId;
_messageId = message.messageId;
}
AIProvider.LAST_ACTION_SESSIONID = _sessionId;
if (isRootSession) {
AIProvider.LAST_ROOT_SESSION_ID = _sessionId;
} }
AIProvider.LAST_ACTION_SESSIONID = sessionId;
return client.chatText({ return client.chatText({
sessionId: _sessionId, sessionId,
messageId: _messageId, messageId,
}); });
})(), })(),
]); ]);
@@ -248,50 +199,36 @@ export function textToText({
// Only one image is currently being processed // Only one image is currently being processed
export function toImage({ export function toImage({
docId,
workspaceId,
promptName,
content, content,
sessionId,
attachments, attachments,
params, params,
seed, seed,
sessionId,
signal, signal,
timeout = TIMEOUT, timeout = TIMEOUT,
retry = false, retry = false,
workflow = false, workflow = false,
client, client,
}: ToImageOptions) { }: ToImageOptions) {
let _sessionId: string; let messageId: string | undefined;
let _messageId: string | undefined;
return { return {
[Symbol.asyncIterator]: async function* () { [Symbol.asyncIterator]: async function* () {
if (retry) { if (!retry) {
const retrySessionId = messageId = await createMessage({
(await sessionId) ?? AIProvider.LAST_ACTION_SESSIONID; client,
_sessionId = retrySessionId; sessionId,
_messageId = undefined;
} else {
const { messageId, sessionId } = await createSessionMessage({
docId,
workspaceId,
promptName,
content, content,
attachments, attachments,
params, params,
client,
}); });
_sessionId = sessionId;
_messageId = messageId;
} }
const eventSource = client.imagesStream( const eventSource = client.imagesStream(
_sessionId, sessionId,
_messageId, messageId,
seed, seed,
workflow ? 'workflow' : undefined workflow ? 'workflow' : undefined
); );
AIProvider.LAST_ACTION_SESSIONID = _sessionId; AIProvider.LAST_ACTION_SESSIONID = sessionId;
for await (const event of toTextStream(eventSource, { for await (const event of toTextStream(eventSource, {
timeout, timeout,

View File

@@ -14,7 +14,7 @@ import type { PromptKey } from './prompt';
import { textToText, toImage } from './request'; import { textToText, toImage } from './request';
import { setupTracker } from './tracker'; import { setupTracker } from './tracker';
const filterStyleToPromptName = new Map( const filterStyleToPromptName = new Map<string, PromptKey>(
Object.entries({ Object.entries({
'Clay style': 'workflow:image-clay', 'Clay style': 'workflow:image-clay',
'Pixel style': 'workflow:image-pixel', 'Pixel style': 'workflow:image-pixel',
@@ -23,7 +23,7 @@ const filterStyleToPromptName = new Map(
}) })
); );
const processTypeToPromptName = new Map( const processTypeToPromptName = new Map<string, PromptKey>(
Object.entries({ Object.entries({
Clearer: 'debug:action:fal-upscaler', Clearer: 'debug:action:fal-upscaler',
'Remove background': 'debug:action:fal-remove-bg', 'Remove background': 'debug:action:fal-remove-bg',
@@ -35,31 +35,78 @@ export function setupAIProvider(
client: CopilotClient, client: CopilotClient,
globalDialogService: GlobalDialogService globalDialogService: GlobalDialogService
) { ) {
async function createSession({
workspaceId,
docId,
promptName,
sessionId,
retry,
}: {
workspaceId: string;
docId: string;
promptName: PromptKey;
sessionId?: string;
retry?: boolean;
}) {
if (sessionId) return sessionId;
if (retry) return AIProvider.LAST_ACTION_SESSIONID;
return client.createSession({
workspaceId,
docId,
promptName,
});
}
//#region actions //#region actions
AIProvider.provide('chat', options => { AIProvider.provide('chat', async options => {
const { input, contexts, ...rest } = options; const { input, contexts, attachments, networkSearch, retry } = options;
const disableSearch =
!!contexts?.files.length ||
!!contexts?.docs.length ||
!!attachments?.length;
const promptName =
networkSearch && !disableSearch
? 'Search With AFFiNE AI'
: 'Chat With AFFiNE AI';
const sessionId = await createSession({
promptName,
...options,
});
if (!retry) {
await AIProvider.session?.updateSession(sessionId, promptName);
}
return textToText({ return textToText({
...rest, ...options,
client, client,
sessionId,
content: input, content: input,
params: contexts, params: contexts,
}); });
}); });
AIProvider.provide('summary', options => { AIProvider.provide('summary', async options => {
const sessionId = await createSession({
promptName: 'Summary',
...options,
});
return textToText({ return textToText({
...options, ...options,
client, client,
sessionId,
content: options.input, content: options.input,
promptName: 'Summary',
}); });
}); });
AIProvider.provide('translate', options => { AIProvider.provide('translate', async options => {
const sessionId = await createSession({
promptName: 'Translate to',
...options,
});
return textToText({ return textToText({
...options, ...options,
client, client,
promptName: 'Translate to', sessionId,
content: options.input, content: options.input,
params: { params: {
language: options.lang, language: options.lang,
@@ -67,200 +114,280 @@ export function setupAIProvider(
}); });
}); });
AIProvider.provide('changeTone', options => { AIProvider.provide('changeTone', async options => {
const sessionId = await createSession({
promptName: 'Change tone to',
...options,
});
return textToText({ return textToText({
...options, ...options,
client, client,
sessionId,
params: { params: {
tone: options.tone.toLowerCase(), tone: options.tone.toLowerCase(),
}, },
content: options.input, content: options.input,
promptName: 'Change tone to',
}); });
}); });
AIProvider.provide('improveWriting', options => { AIProvider.provide('improveWriting', async options => {
return textToText({ const sessionId = await createSession({
...options,
client,
content: options.input,
promptName: 'Improve writing for it', promptName: 'Improve writing for it',
...options,
}); });
});
AIProvider.provide('improveGrammar', options => {
return textToText({ return textToText({
...options, ...options,
client, client,
sessionId,
content: options.input, content: options.input,
});
});
AIProvider.provide('improveGrammar', async options => {
const sessionId = await createSession({
promptName: 'Improve grammar for it', promptName: 'Improve grammar for it',
...options,
}); });
});
AIProvider.provide('fixSpelling', options => {
return textToText({ return textToText({
...options, ...options,
client, client,
sessionId,
content: options.input, content: options.input,
});
});
AIProvider.provide('fixSpelling', async options => {
const sessionId = await createSession({
promptName: 'Fix spelling for it', promptName: 'Fix spelling for it',
...options,
}); });
});
AIProvider.provide('createHeadings', options => {
return textToText({ return textToText({
...options, ...options,
client, client,
sessionId,
content: options.input, content: options.input,
});
});
AIProvider.provide('createHeadings', async options => {
const sessionId = await createSession({
promptName: 'Create headings', promptName: 'Create headings',
...options,
}); });
});
AIProvider.provide('makeLonger', options => {
return textToText({ return textToText({
...options, ...options,
client, client,
sessionId,
content: options.input, content: options.input,
});
});
AIProvider.provide('makeLonger', async options => {
const sessionId = await createSession({
promptName: 'Make it longer', promptName: 'Make it longer',
...options,
}); });
});
AIProvider.provide('makeShorter', options => {
return textToText({ return textToText({
...options, ...options,
client, client,
sessionId,
content: options.input, content: options.input,
});
});
AIProvider.provide('makeShorter', async options => {
const sessionId = await createSession({
promptName: 'Make it shorter', promptName: 'Make it shorter',
...options,
}); });
});
AIProvider.provide('checkCodeErrors', options => {
return textToText({ return textToText({
...options, ...options,
client, client,
sessionId,
content: options.input, content: options.input,
});
});
AIProvider.provide('checkCodeErrors', async options => {
const sessionId = await createSession({
promptName: 'Check code error', promptName: 'Check code error',
...options,
}); });
});
AIProvider.provide('explainCode', options => {
return textToText({ return textToText({
...options, ...options,
client, client,
sessionId,
content: options.input, content: options.input,
});
});
AIProvider.provide('explainCode', async options => {
const sessionId = await createSession({
promptName: 'Explain this code', promptName: 'Explain this code',
...options,
}); });
});
AIProvider.provide('writeArticle', options => {
return textToText({ return textToText({
...options, ...options,
client, client,
sessionId,
content: options.input, content: options.input,
});
});
AIProvider.provide('writeArticle', async options => {
const sessionId = await createSession({
promptName: 'Write an article about this', promptName: 'Write an article about this',
...options,
}); });
});
AIProvider.provide('writeTwitterPost', options => {
return textToText({ return textToText({
...options, ...options,
client, client,
sessionId,
content: options.input, content: options.input,
});
});
AIProvider.provide('writeTwitterPost', async options => {
const sessionId = await createSession({
promptName: 'Write a twitter about this', promptName: 'Write a twitter about this',
...options,
}); });
});
AIProvider.provide('writePoem', options => {
return textToText({ return textToText({
...options, ...options,
client, client,
sessionId,
content: options.input, content: options.input,
});
});
AIProvider.provide('writePoem', async options => {
const sessionId = await createSession({
promptName: 'Write a poem about this', promptName: 'Write a poem about this',
...options,
}); });
});
AIProvider.provide('writeOutline', options => {
return textToText({ return textToText({
...options, ...options,
client, client,
sessionId,
content: options.input, content: options.input,
});
});
AIProvider.provide('writeOutline', async options => {
const sessionId = await createSession({
promptName: 'Write outline', promptName: 'Write outline',
...options,
}); });
});
AIProvider.provide('writeBlogPost', options => {
return textToText({ return textToText({
...options, ...options,
client, client,
sessionId,
content: options.input, content: options.input,
});
});
AIProvider.provide('writeBlogPost', async options => {
const sessionId = await createSession({
promptName: 'Write a blog post about this', promptName: 'Write a blog post about this',
...options,
}); });
});
AIProvider.provide('brainstorm', options => {
return textToText({ return textToText({
...options, ...options,
client, client,
sessionId,
content: options.input, content: options.input,
});
});
AIProvider.provide('brainstorm', async options => {
const sessionId = await createSession({
promptName: 'Brainstorm ideas about this', promptName: 'Brainstorm ideas about this',
...options,
}); });
});
AIProvider.provide('findActions', options => {
return textToText({ return textToText({
...options, ...options,
client, client,
sessionId,
content: options.input, content: options.input,
});
});
AIProvider.provide('findActions', async options => {
const sessionId = await createSession({
promptName: 'Find action items from it', promptName: 'Find action items from it',
...options,
}); });
});
AIProvider.provide('brainstormMindmap', options => {
return textToText({ return textToText({
...options, ...options,
client, client,
sessionId,
content: options.input, content: options.input,
});
});
AIProvider.provide('brainstormMindmap', async options => {
const sessionId = await createSession({
promptName: 'workflow:brainstorm', promptName: 'workflow:brainstorm',
...options,
});
return textToText({
...options,
client,
sessionId,
content: options.input,
// 3 minutes // 3 minutes
timeout: 180000, timeout: 180000,
workflow: true, workflow: true,
}); });
}); });
AIProvider.provide('expandMindmap', options => { AIProvider.provide('expandMindmap', async options => {
if (!options.input) { if (!options.input) {
throw new Error('expandMindmap action requires input'); throw new Error('expandMindmap action requires input');
} }
const sessionId = await createSession({
promptName: 'Expand mind map',
...options,
});
return textToText({ return textToText({
...options, ...options,
client, client,
sessionId,
params: { params: {
mindmap: options.mindmap, mindmap: options.mindmap,
node: options.input, node: options.input,
}, },
content: options.input, content: options.input,
promptName: 'Expand mind map',
}); });
}); });
AIProvider.provide('explain', options => { AIProvider.provide('explain', async options => {
return textToText({ const sessionId = await createSession({
...options,
client,
content: options.input,
promptName: 'Explain this', promptName: 'Explain this',
...options,
}); });
});
AIProvider.provide('explainImage', options => {
return textToText({ return textToText({
...options, ...options,
client, client,
sessionId,
content: options.input, content: options.input,
promptName: 'Explain this image',
}); });
}); });
AIProvider.provide('makeItReal', options => { AIProvider.provide('explainImage', async options => {
const sessionId = await createSession({
promptName: 'Explain this image',
...options,
});
return textToText({
...options,
client,
sessionId,
content: options.input,
});
});
AIProvider.provide('makeItReal', async options => {
let promptName: PromptKey = 'Make it real'; let promptName: PromptKey = 'Make it real';
let content = options.input || ''; let content = options.input || '';
@@ -275,15 +402,20 @@ Here are our design notes:\n ${content}.`;
Could you make a new website based on these notes and send back just the html file?`; Could you make a new website based on these notes and send back just the html file?`;
} }
const sessionId = await createSession({
promptName,
...options,
});
return textToText({ return textToText({
...options, ...options,
client, client,
sessionId,
content, content,
promptName,
}); });
}); });
AIProvider.provide('createSlides', options => { AIProvider.provide('createSlides', async options => {
const SlideSchema = z.object({ const SlideSchema = z.object({
page: z.number(), page: z.number(),
type: z.enum(['name', 'title', 'content']), type: z.enum(['name', 'title', 'content']),
@@ -320,11 +452,15 @@ Could you make a new website based on these notes and send back just the html fi
}) })
.join('\n'); .join('\n');
}; };
const sessionId = await createSession({
promptName: 'workflow:presentation',
...options,
});
return textToText({ return textToText({
...options, ...options,
client, client,
sessionId,
content: options.input, content: options.input,
promptName: 'workflow:presentation',
// 3 minutes // 3 minutes
timeout: 180000, timeout: 180000,
workflow: true, workflow: true,
@@ -332,79 +468,98 @@ Could you make a new website based on these notes and send back just the html fi
}); });
}); });
AIProvider.provide('createImage', options => { AIProvider.provide('createImage', async options => {
// test to image // test to image
let promptName: PromptKey = 'debug:action:dalle3'; let promptName: PromptKey = 'debug:action:dalle3';
// image to image // image to image
if (options.attachments?.length) { if (options.attachments?.length) {
promptName = 'debug:action:fal-sd15'; promptName = 'debug:action:fal-sd15';
} }
const sessionId = await createSession({
promptName,
...options,
});
return toImage({ return toImage({
...options, ...options,
client, client,
sessionId,
content: options.input, content: options.input,
promptName,
}); });
}); });
AIProvider.provide('filterImage', options => { AIProvider.provide('filterImage', async options => {
// test to image // test to image
const promptName = filterStyleToPromptName.get(options.style as string); const promptName: PromptKey | undefined = filterStyleToPromptName.get(
options.style
);
if (!promptName) {
throw new Error('filterImage requires a promptName');
}
const sessionId = await createSession({
promptName,
...options,
});
return toImage({ return toImage({
...options, ...options,
client, client,
sessionId,
content: options.input, content: options.input,
timeout: 180000, timeout: 180000,
promptName: promptName as PromptKey,
workflow: !!promptName?.startsWith('workflow:'), workflow: !!promptName?.startsWith('workflow:'),
}); });
}); });
AIProvider.provide('processImage', options => { AIProvider.provide('processImage', async options => {
// test to image // test to image
const promptName = processTypeToPromptName.get( const promptName: PromptKey | undefined = processTypeToPromptName.get(
options.type as string options.type
) as PromptKey; );
if (!promptName) {
throw new Error('processImage requires a promptName');
}
const sessionId = await createSession({
promptName,
...options,
});
return toImage({ return toImage({
...options, ...options,
client, client,
sessionId,
content: options.input, content: options.input,
timeout: 180000, timeout: 180000,
promptName,
}); });
}); });
AIProvider.provide('generateCaption', options => { AIProvider.provide('generateCaption', async options => {
return textToText({ const sessionId = await createSession({
...options,
client,
content: options.input,
promptName: 'Generate a caption', promptName: 'Generate a caption',
...options,
}); });
});
AIProvider.provide('continueWriting', options => {
return textToText({ return textToText({
...options, ...options,
client, client,
sessionId,
content: options.input, content: options.input,
});
});
AIProvider.provide('continueWriting', async options => {
const sessionId = await createSession({
promptName: 'Continue writing', promptName: 'Continue writing',
...options,
});
return textToText({
...options,
client,
sessionId,
content: options.input,
}); });
}); });
//#endregion //#endregion
AIProvider.provide('session', { AIProvider.provide('session', {
createSession: async ( createSession,
workspaceId: string,
docId: string,
promptName = 'Chat With AFFiNE AI'
) => {
return client.createSession({
workspaceId,
docId,
promptName,
});
},
getSessions: async ( getSessions: async (
workspaceId: string, workspaceId: string,
docId?: string, docId?: string,