mirror of
https://github.com/toeverything/AFFiNE.git
synced 2026-07-25 06:18:45 +08:00
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:
@@ -20,9 +20,10 @@ import {
|
||||
import { MockEmbeddingClient } from '../plugins/copilot/context/embedding';
|
||||
import { prompts, PromptService } from '../plugins/copilot/prompt';
|
||||
import {
|
||||
CopilotCapability,
|
||||
CopilotProviderFactory,
|
||||
CopilotProviderType,
|
||||
ModelInputType,
|
||||
ModelOutputType,
|
||||
OpenAIProvider,
|
||||
} from '../plugins/copilot/providers';
|
||||
import { CitationParser } from '../plugins/copilot/providers/utils';
|
||||
@@ -756,9 +757,7 @@ test('should be able to get provider', async t => {
|
||||
const { factory } = t.context;
|
||||
|
||||
{
|
||||
const p = await factory.getProviderByCapability(
|
||||
CopilotCapability.TextToText
|
||||
);
|
||||
const p = await factory.getProvider({ outputType: ModelOutputType.Text });
|
||||
t.is(
|
||||
p?.type.toString(),
|
||||
'openai',
|
||||
@@ -767,36 +766,41 @@ test('should be able to get provider', async t => {
|
||||
}
|
||||
|
||||
{
|
||||
const p = await factory.getProviderByCapability(
|
||||
CopilotCapability.ImageToImage,
|
||||
{ model: 'lora/image-to-image' }
|
||||
);
|
||||
const p = await factory.getProvider({
|
||||
outputType: ModelOutputType.Image,
|
||||
inputTypes: [ModelInputType.Image],
|
||||
modelId: 'lora/image-to-image',
|
||||
});
|
||||
t.is(
|
||||
p?.type.toString(),
|
||||
'fal',
|
||||
'should get provider support text-to-embedding'
|
||||
'should get provider supporting image output'
|
||||
);
|
||||
}
|
||||
|
||||
{
|
||||
const p = await factory.getProviderByCapability(
|
||||
CopilotCapability.ImageToText,
|
||||
const p = await factory.getProvider(
|
||||
{
|
||||
outputType: ModelOutputType.Image,
|
||||
inputTypes: [ModelInputType.Image],
|
||||
},
|
||||
{ prefer: CopilotProviderType.FAL }
|
||||
);
|
||||
t.is(
|
||||
p?.type.toString(),
|
||||
'fal',
|
||||
'should get provider support text-to-embedding'
|
||||
'should get provider supporting text output with image input'
|
||||
);
|
||||
}
|
||||
|
||||
// if a model is not defined and not available in online api
|
||||
// it should return null
|
||||
{
|
||||
const p = await factory.getProviderByCapability(
|
||||
CopilotCapability.ImageToText,
|
||||
{ model: 'gpt-4-not-exist' }
|
||||
);
|
||||
const p = await factory.getProvider({
|
||||
outputType: ModelOutputType.Text,
|
||||
inputTypes: [ModelInputType.Text],
|
||||
modelId: 'gpt-4-not-exist',
|
||||
});
|
||||
t.falsy(p, 'should not get provider');
|
||||
}
|
||||
});
|
||||
@@ -987,10 +991,9 @@ test('should be able to run text executor', async t => {
|
||||
{ role: 'system', content: 'hello {{word}}' },
|
||||
]);
|
||||
// mock provider
|
||||
const testProvider =
|
||||
(await factory.getProviderByModel<CopilotCapability.TextToText>('test'))!;
|
||||
const text = Sinon.spy(testProvider, 'generateText');
|
||||
const textStream = Sinon.spy(testProvider, 'generateTextStream');
|
||||
const testProvider = (await factory.getProviderByModel('test'))!;
|
||||
const text = Sinon.spy(testProvider, 'text');
|
||||
const textStream = Sinon.spy(testProvider, 'streamText');
|
||||
|
||||
const nodeData: WorkflowNodeData = {
|
||||
id: 'basic',
|
||||
@@ -1013,7 +1016,7 @@ test('should be able to run text executor', async t => {
|
||||
},
|
||||
]);
|
||||
t.deepEqual(
|
||||
text.lastCall.args[0][0].content,
|
||||
text.lastCall.args[1][0].content,
|
||||
'hello world',
|
||||
'should render the prompt with params'
|
||||
);
|
||||
@@ -1036,7 +1039,7 @@ test('should be able to run text executor', async t => {
|
||||
}))
|
||||
);
|
||||
t.deepEqual(
|
||||
textStream.lastCall.args[0][0].params?.attachments,
|
||||
textStream.lastCall.args[1][0].params?.attachments,
|
||||
['https://affine.pro/example.jpg'],
|
||||
'should pass attachments to provider'
|
||||
);
|
||||
@@ -1050,14 +1053,13 @@ test('should be able to run image executor', async t => {
|
||||
|
||||
executors.image.register();
|
||||
const executor = getWorkflowExecutor(executors.image.type);
|
||||
await prompt.set('test', 'test', [
|
||||
await prompt.set('test', 'test-image', [
|
||||
{ role: 'user', content: 'tag1, tag2, tag3, {{#tags}}{{.}}, {{/tags}}' },
|
||||
]);
|
||||
// mock provider
|
||||
const testProvider =
|
||||
(await factory.getProviderByModel<CopilotCapability.TextToImage>('test'))!;
|
||||
const image = Sinon.spy(testProvider, 'generateImages');
|
||||
const imageStream = Sinon.spy(testProvider, 'generateImagesStream');
|
||||
const testProvider = (await factory.getProviderByModel('test'))!;
|
||||
|
||||
const imageStream = Sinon.spy(testProvider, 'streamImages');
|
||||
|
||||
const nodeData: WorkflowNodeData = {
|
||||
id: 'basic',
|
||||
@@ -1076,20 +1078,9 @@ test('should be able to run image executor', async t => {
|
||||
)
|
||||
);
|
||||
|
||||
t.deepEqual(ret, [
|
||||
{
|
||||
type: NodeExecuteState.Params,
|
||||
params: {
|
||||
key: [
|
||||
'https://example.com/test.jpg',
|
||||
'tag1, tag2, tag3, tag4, tag5, ',
|
||||
],
|
||||
},
|
||||
},
|
||||
]);
|
||||
t.deepEqual(
|
||||
image.lastCall.args[0][0].content,
|
||||
'tag1, tag2, tag3, tag4, tag5, ',
|
||||
t.snapshot(ret, 'should generate image stream');
|
||||
t.snapshot(
|
||||
imageStream.lastCall.args,
|
||||
'should render the prompt with params array'
|
||||
);
|
||||
}
|
||||
@@ -1104,16 +1095,17 @@ test('should be able to run image executor', async t => {
|
||||
|
||||
t.deepEqual(
|
||||
ret,
|
||||
Array.from(['https://example.com/test.jpg', 'tag1, tag2, tag3, ']).map(
|
||||
t => ({
|
||||
attachment: t,
|
||||
nodeId: 'basic',
|
||||
type: NodeExecuteState.Attachment,
|
||||
})
|
||||
)
|
||||
Array.from([
|
||||
'https://example.com/test-image.jpg',
|
||||
'tag1, tag2, tag3, ',
|
||||
]).map(t => ({
|
||||
attachment: t,
|
||||
nodeId: 'basic',
|
||||
type: NodeExecuteState.Attachment,
|
||||
}))
|
||||
);
|
||||
t.deepEqual(
|
||||
imageStream.lastCall.args[0][0].params?.attachments,
|
||||
imageStream.lastCall.args[1][0].params?.attachments,
|
||||
['https://affine.pro/example.jpg'],
|
||||
'should pass attachments to provider'
|
||||
);
|
||||
|
||||
Reference in New Issue
Block a user