feat: add prompt level config (#7445)

This commit is contained in:
darkskygit
2024-07-08 08:11:22 +00:00
parent 9ef8829ef1
commit bf6c9a5955
12 changed files with 125 additions and 41 deletions

View File

@@ -0,0 +1,2 @@
-- AlterTable
ALTER TABLE "ai_prompts_metadata" ADD COLUMN "config" JSON;

View File

@@ -457,6 +457,7 @@ model AiPrompt {
// it is only used in the frontend and does not affect the backend // it is only used in the frontend and does not affect the backend
action String? @db.VarChar action String? @db.VarChar
model String @db.VarChar model String @db.VarChar
config Json? @db.Json
createdAt DateTime @default(now()) @map("created_at") @db.Timestamptz(6) createdAt DateTime @default(now()) @map("created_at") @db.Timestamptz(6)
messages AiPromptMessage[] messages AiPromptMessage[]

View File

@@ -6,10 +6,19 @@ type PromptMessage = {
params?: Record<string, string | string[]>; params?: Record<string, string | string[]>;
}; };
type PromptConfig = {
jsonMode?: boolean;
frequencyPenalty?: number;
presencePenalty?: number;
temperature?: number;
maxTokens?: number;
};
type Prompt = { type Prompt = {
name: string; name: string;
action?: string; action?: string;
model: string; model: string;
config?: PromptConfig;
messages: PromptMessage[]; messages: PromptMessage[];
}; };
@@ -465,6 +474,7 @@ content: {{content}}`,
name: 'workflow:presentation:step1', name: 'workflow:presentation:step1',
action: 'workflow:presentation:step1', action: 'workflow:presentation:step1',
model: 'gpt-4o', model: 'gpt-4o',
config: { temperature: 0.7 },
messages: [ messages: [
{ {
role: 'system', role: 'system',
@@ -685,6 +695,7 @@ export async function refreshPrompts(db: PrismaClient) {
create: { create: {
name: prompt.name, name: prompt.name,
action: prompt.action, action: prompt.action,
config: prompt.config,
model: prompt.model, model: prompt.model,
messages: { messages: {
create: prompt.messages.map((message, idx) => ({ create: prompt.messages.map((message, idx) => ({

View File

@@ -138,9 +138,8 @@ export class CopilotController {
const messageId = Array.isArray(params.messageId) const messageId = Array.isArray(params.messageId)
? params.messageId[0] ? params.messageId[0]
: params.messageId; : params.messageId;
const jsonMode = String(params.jsonMode).toLowerCase() === 'true';
delete params.messageId; delete params.messageId;
return { messageId, jsonMode, params }; return { messageId, params };
} }
private getSignal(req: Request) { private getSignal(req: Request) {
@@ -167,7 +166,7 @@ export class CopilotController {
@Param('sessionId') sessionId: string, @Param('sessionId') sessionId: string,
@Query() params: Record<string, string | string[]> @Query() params: Record<string, string | string[]>
): Promise<string> { ): Promise<string> {
const { messageId, jsonMode } = this.prepareParams(params); const { messageId } = this.prepareParams(params);
const provider = await this.chooseTextProvider( const provider = await this.chooseTextProvider(
user.id, user.id,
sessionId, sessionId,
@@ -180,7 +179,11 @@ export class CopilotController {
const content = await provider.generateText( const content = await provider.generateText(
session.finish(params), session.finish(params),
session.model, session.model,
{ jsonMode, signal: this.getSignal(req), user: user.id } {
...session.config.promptConfig,
signal: this.getSignal(req),
user: user.id,
}
); );
session.push({ session.push({
@@ -204,7 +207,7 @@ export class CopilotController {
@Query() params: Record<string, string> @Query() params: Record<string, string>
): Promise<Observable<ChatEvent>> { ): Promise<Observable<ChatEvent>> {
try { try {
const { messageId, jsonMode } = this.prepareParams(params); const { messageId } = this.prepareParams(params);
const provider = await this.chooseTextProvider( const provider = await this.chooseTextProvider(
user.id, user.id,
sessionId, sessionId,
@@ -215,7 +218,7 @@ export class CopilotController {
return from( return from(
provider.generateTextStream(session.finish(params), session.model, { provider.generateTextStream(session.finish(params), session.model, {
jsonMode, ...session.config.promptConfig,
signal: this.getSignal(req), signal: this.getSignal(req),
user: user.id, user: user.id,
}) })
@@ -256,7 +259,7 @@ export class CopilotController {
@Query() params: Record<string, string> @Query() params: Record<string, string>
): Promise<Observable<ChatEvent>> { ): Promise<Observable<ChatEvent>> {
try { try {
const { messageId, jsonMode } = this.prepareParams(params); const { messageId } = this.prepareParams(params);
const session = await this.appendSessionMessage(sessionId, messageId); const session = await this.appendSessionMessage(sessionId, messageId);
const latestMessage = session.stashMessages.findLast( const latestMessage = session.stashMessages.findLast(
m => m.role === 'user' m => m.role === 'user'
@@ -269,7 +272,7 @@ export class CopilotController {
return from( return from(
this.workflow.runGraph(params, session.model, { this.workflow.runGraph(params, session.model, {
jsonMode, ...session.config.promptConfig,
signal: this.getSignal(req), signal: this.getSignal(req),
user: user.id, user: user.id,
}) })

View File

@@ -5,6 +5,8 @@ import Mustache from 'mustache';
import { import {
getTokenEncoder, getTokenEncoder,
PromptConfig,
PromptConfigSchema,
PromptMessage, PromptMessage,
PromptMessageSchema, PromptMessageSchema,
PromptParams, PromptParams,
@@ -35,14 +37,16 @@ export class ChatPrompt {
private readonly templateParams: PromptParams = {}; private readonly templateParams: PromptParams = {};
static createFromPrompt( static createFromPrompt(
options: Omit<AiPrompt, 'id' | 'createdAt'> & { options: Omit<AiPrompt, 'id' | 'createdAt' | 'config'> & {
messages: PromptMessage[]; messages: PromptMessage[];
config: PromptConfig | undefined;
} }
) { ) {
return new ChatPrompt( return new ChatPrompt(
options.name, options.name,
options.action || undefined, options.action || undefined,
options.model, options.model,
options.config,
options.messages options.messages
); );
} }
@@ -51,6 +55,7 @@ export class ChatPrompt {
public readonly name: string, public readonly name: string,
public readonly action: string | undefined, public readonly action: string | undefined,
public readonly model: string, public readonly model: string,
public readonly config: PromptConfig | undefined,
private readonly messages: PromptMessage[] private readonly messages: PromptMessage[]
) { ) {
this.encoder = getTokenEncoder(model); this.encoder = getTokenEncoder(model);
@@ -185,6 +190,7 @@ export class PromptService {
name: true, name: true,
action: true, action: true,
model: true, model: true,
config: true,
messages: { messages: {
select: { select: {
role: true, role: true,
@@ -199,9 +205,11 @@ export class PromptService {
}); });
const messages = PromptMessageSchema.array().safeParse(prompt?.messages); const messages = PromptMessageSchema.array().safeParse(prompt?.messages);
if (prompt && messages.success) { const config = PromptConfigSchema.safeParse(prompt?.config);
if (prompt && messages.success && config.success) {
const chatPrompt = ChatPrompt.createFromPrompt({ const chatPrompt = ChatPrompt.createFromPrompt({
...prompt, ...prompt,
config: config.data,
messages: messages.data, messages: messages.data,
}); });
this.cache.set(name, chatPrompt); this.cache.set(name, chatPrompt);
@@ -210,12 +218,18 @@ export class PromptService {
return null; return null;
} }
async set(name: string, model: string, messages: PromptMessage[]) { async set(
name: string,
model: string,
messages: PromptMessage[],
config?: PromptConfig
) {
return await this.db.aiPrompt return await this.db.aiPrompt
.create({ .create({
data: { data: {
name, name,
model, model,
config: config || undefined,
messages: { messages: {
create: messages.map((m, idx) => ({ create: messages.map((m, idx) => ({
idx, idx,
@@ -229,10 +243,11 @@ export class PromptService {
.then(ret => ret.id); .then(ret => ret.id);
} }
async update(name: string, messages: PromptMessage[]) { async update(name: string, messages: PromptMessage[], config?: PromptConfig) {
const { id } = await this.db.aiPrompt.update({ const { id } = await this.db.aiPrompt.update({
where: { name }, where: { name },
data: { data: {
config: config || undefined,
messages: { messages: {
// cleanup old messages // cleanup old messages
deleteMany: {}, deleteMany: {},

View File

@@ -125,21 +125,6 @@ export class OpenAIProvider
}); });
} }
private extractOptionFromMessages(
messages: PromptMessage[],
options: CopilotChatOptions
) {
const params: Record<string, string | string[]> = {};
for (const message of messages) {
if (message.params) {
Object.assign(params, message.params);
}
}
if (params.jsonMode && options) {
options.jsonMode = String(params.jsonMode).toLowerCase() === 'true';
}
}
protected checkParams({ protected checkParams({
messages, messages,
embeddings, embeddings,
@@ -155,7 +140,6 @@ export class OpenAIProvider
throw new CopilotPromptInvalid(`Invalid model: ${model}`); throw new CopilotPromptInvalid(`Invalid model: ${model}`);
} }
if (Array.isArray(messages) && messages.length > 0) { if (Array.isArray(messages) && messages.length > 0) {
this.extractOptionFromMessages(messages, options);
if ( if (
messages.some( messages.some(
m => m =>
@@ -257,7 +241,9 @@ export class OpenAIProvider
stream: true, stream: true,
messages: this.chatToGPTMessage(messages), messages: this.chatToGPTMessage(messages),
model: model, model: model,
temperature: options.temperature || 0, frequency_penalty: options.frequencyPenalty || 0,
presence_penalty: options.presencePenalty || 0,
temperature: options.temperature || 0.5,
max_tokens: options.maxTokens || 4096, max_tokens: options.maxTokens || 4096,
response_format: { response_format: {
type: options.jsonMode ? 'json_object' : 'text', type: options.jsonMode ? 'json_object' : 'text',

View File

@@ -183,6 +183,25 @@ registerEnumType(AiPromptRole, {
name: 'CopilotPromptMessageRole', name: 'CopilotPromptMessageRole',
}); });
@InputType('CopilotPromptConfigInput')
@ObjectType()
class CopilotPromptConfigType {
@Field(() => Boolean, { nullable: true })
jsonMode!: boolean | null;
@Field(() => Number, { nullable: true })
frequencyPenalty!: number | null;
@Field(() => Number, { nullable: true })
presencePenalty!: number | null;
@Field(() => Number, { nullable: true })
temperature!: number | null;
@Field(() => Number, { nullable: true })
topP!: number | null;
}
@InputType('CopilotPromptMessageInput') @InputType('CopilotPromptMessageInput')
@ObjectType() @ObjectType()
class CopilotPromptMessageType { class CopilotPromptMessageType {
@@ -209,6 +228,9 @@ class CopilotPromptType {
@Field(() => String, { nullable: true }) @Field(() => String, { nullable: true })
action!: string | null; action!: string | null;
@Field(() => CopilotPromptConfigType, { nullable: true })
config!: CopilotPromptConfigType | null;
@Field(() => [CopilotPromptMessageType]) @Field(() => [CopilotPromptMessageType])
messages!: CopilotPromptMessageType[]; messages!: CopilotPromptMessageType[];
} }
@@ -462,6 +484,9 @@ class CreateCopilotPromptInput {
@Field(() => String, { nullable: true }) @Field(() => String, { nullable: true })
action!: string | null; action!: string | null;
@Field(() => CopilotPromptConfigType, { nullable: true })
config!: CopilotPromptConfigType | null;
@Field(() => [CopilotPromptMessageType]) @Field(() => [CopilotPromptMessageType])
messages!: CopilotPromptMessageType[]; messages!: CopilotPromptMessageType[];
} }
@@ -485,7 +510,12 @@ export class PromptsManagementResolver {
@Args({ type: () => CreateCopilotPromptInput, name: 'input' }) @Args({ type: () => CreateCopilotPromptInput, name: 'input' })
input: CreateCopilotPromptInput input: CreateCopilotPromptInput
) { ) {
await this.promptService.set(input.name, input.model, input.messages); await this.promptService.set(
input.name,
input.model,
input.messages,
input.config
);
return this.promptService.get(input.name); return this.promptService.get(input.name);
} }

View File

@@ -49,10 +49,10 @@ export class ChatSession implements AsyncDisposable {
userId, userId,
workspaceId, workspaceId,
docId, docId,
prompt: { name: promptName }, prompt: { name: promptName, config: promptConfig },
} = this.state; } = this.state;
return { sessionId, userId, workspaceId, docId, promptName }; return { sessionId, userId, workspaceId, docId, promptName, promptConfig };
} }
get stashMessages() { get stashMessages() {

View File

@@ -63,6 +63,20 @@ export type PromptMessage = z.infer<typeof PromptMessageSchema>;
export type PromptParams = NonNullable<PromptMessage['params']>; export type PromptParams = NonNullable<PromptMessage['params']>;
export const PromptConfigStrictSchema = z.object({
jsonMode: z.boolean().nullable().optional(),
frequencyPenalty: z.number().nullable().optional(),
presencePenalty: z.number().nullable().optional(),
temperature: z.number().nullable().optional(),
topP: z.number().nullable().optional(),
maxTokens: z.number().nullable().optional(),
});
export const PromptConfigSchema =
PromptConfigStrictSchema.nullable().optional();
export type PromptConfig = z.infer<typeof PromptConfigSchema>;
export const ChatMessageSchema = PromptMessageSchema.extend({ export const ChatMessageSchema = PromptMessageSchema.extend({
id: z.string().optional(), id: z.string().optional(),
createdAt: z.date(), createdAt: z.date(),
@@ -144,11 +158,9 @@ const CopilotProviderOptionsSchema = z.object({
user: z.string().optional(), user: z.string().optional(),
}); });
const CopilotChatOptionsSchema = CopilotProviderOptionsSchema.extend({ const CopilotChatOptionsSchema = CopilotProviderOptionsSchema.merge(
jsonMode: z.boolean().optional(), PromptConfigStrictSchema
temperature: z.number().optional(), ).optional();
maxTokens: z.number().optional(),
}).optional();
export type CopilotChatOptions = z.infer<typeof CopilotChatOptionsSchema>; export type CopilotChatOptions = z.infer<typeof CopilotChatOptionsSchema>;

View File

@@ -57,6 +57,22 @@ enum CopilotModels {
TextModerationStable TextModerationStable
} }
input CopilotPromptConfigInput {
frequencyPenalty: Int
jsonMode: Boolean
presencePenalty: Int
temperature: Int
topP: Int
}
type CopilotPromptConfigType {
frequencyPenalty: Int
jsonMode: Boolean
presencePenalty: Int
temperature: Int
topP: Int
}
input CopilotPromptMessageInput { input CopilotPromptMessageInput {
content: String! content: String!
params: JSON params: JSON
@@ -81,6 +97,7 @@ type CopilotPromptNotFoundDataType {
type CopilotPromptType { type CopilotPromptType {
action: String action: String
config: CopilotPromptConfigType
messages: [CopilotPromptMessageType!]! messages: [CopilotPromptMessageType!]!
model: CopilotModels! model: CopilotModels!
name: String! name: String!
@@ -123,6 +140,7 @@ input CreateCheckoutSessionInput {
input CreateCopilotPromptInput { input CreateCopilotPromptInput {
action: String action: String
config: CopilotPromptConfigInput
messages: [CopilotPromptMessageInput!]! messages: [CopilotPromptMessageInput!]!
model: CopilotModels! model: CopilotModels!
name: String! name: String!

View File

@@ -676,7 +676,7 @@ test.skip('should be able to preview workflow', async t => {
registerCopilotProvider(OpenAIProvider); registerCopilotProvider(OpenAIProvider);
for (const p of prompts) { for (const p of prompts) {
await prompt.set(p.name, p.model, p.messages); await prompt.set(p.name, p.model, p.messages, p.config);
} }
let result = ''; let result = '';
@@ -726,7 +726,7 @@ test('should be able to run pre defined workflow', async t => {
const { graph, prompts, callCount, input, params, result } = testCase; const { graph, prompts, callCount, input, params, result } = testCase;
console.log('running workflow test:', graph.name); console.log('running workflow test:', graph.name);
for (const p of prompts) { for (const p of prompts) {
await prompt.set(p.name, p.model, p.messages); await prompt.set(p.name, p.model, p.messages, p.config);
} }
for (const [idx, i] of input.entries()) { for (const [idx, i] of input.entries()) {
@@ -773,7 +773,7 @@ test('should be able to run workflow', async t => {
const executor = Sinon.spy(executors.text, 'next'); const executor = Sinon.spy(executors.text, 'next');
for (const p of prompts) { for (const p of prompts) {
await prompt.set(p.name, p.model, p.messages); await prompt.set(p.name, p.model, p.messages, p.config);
} }
const graphName = 'presentation'; const graphName = 'presentation';

View File

@@ -17,6 +17,7 @@ import {
CopilotTextToEmbeddingProvider, CopilotTextToEmbeddingProvider,
CopilotTextToImageProvider, CopilotTextToImageProvider,
CopilotTextToTextProvider, CopilotTextToTextProvider,
PromptConfig,
PromptMessage, PromptMessage,
} from '../../src/plugins/copilot/types'; } from '../../src/plugins/copilot/types';
import { NodeExecutorType } from '../../src/plugins/copilot/workflow/executor'; import { NodeExecutorType } from '../../src/plugins/copilot/workflow/executor';
@@ -383,7 +384,12 @@ export async function getHistories(
return res.body.data.currentUser?.copilot?.histories || []; return res.body.data.currentUser?.copilot?.histories || [];
} }
type Prompt = { name: string; model: string; messages: PromptMessage[] }; type Prompt = {
name: string;
model: string;
messages: PromptMessage[];
config?: PromptConfig;
};
type WorkflowTestCase = { type WorkflowTestCase = {
graph: WorkflowGraph; graph: WorkflowGraph;
prompts: Prompt[]; prompts: Prompt[];