@@ -0,0 +1,2 @@
|
|||||||
|
-- AlterTable
|
||||||
|
ALTER TABLE "ai_sessions_metadata" ADD COLUMN "title" VARCHAR;
|
||||||
@@ -443,6 +443,7 @@ model AiSession {
|
|||||||
promptName String @map("prompt_name") @db.VarChar(32)
|
promptName String @map("prompt_name") @db.VarChar(32)
|
||||||
promptAction String? @default("") @map("prompt_action") @db.VarChar(32)
|
promptAction String? @default("") @map("prompt_action") @db.VarChar(32)
|
||||||
pinned Boolean @default(false)
|
pinned Boolean @default(false)
|
||||||
|
title String? @db.VarChar
|
||||||
// the session id of the parent session if this session is a forked session
|
// the session id of the parent session if this session is a forked session
|
||||||
parentSessionId String? @map("parent_session_id") @db.VarChar
|
parentSessionId String? @map("parent_session_id") @db.VarChar
|
||||||
messageCost Int @default(0)
|
messageCost Int @default(0)
|
||||||
|
|||||||
@@ -330,3 +330,45 @@ Generated by [AVA](https://avajs.dev).
|
|||||||
],
|
],
|
||||||
{},
|
{},
|
||||||
]
|
]
|
||||||
|
|
||||||
|
## should handle generateSessionTitle correctly under various conditions
|
||||||
|
|
||||||
|
> should generate title when conditions are met
|
||||||
|
|
||||||
|
{
|
||||||
|
chatWithPromptCalled: undefined,
|
||||||
|
exists: true,
|
||||||
|
title: 'What is Machine Learning?',
|
||||||
|
}
|
||||||
|
|
||||||
|
> should not generate title when session already has title
|
||||||
|
|
||||||
|
{
|
||||||
|
chatWithPromptCalled: false,
|
||||||
|
exists: true,
|
||||||
|
title: 'Existing Title',
|
||||||
|
}
|
||||||
|
|
||||||
|
> should not generate title when no user messages exist
|
||||||
|
|
||||||
|
{
|
||||||
|
chatWithPromptCalled: false,
|
||||||
|
exists: true,
|
||||||
|
title: null,
|
||||||
|
}
|
||||||
|
|
||||||
|
> should not generate title when no assistant messages exist
|
||||||
|
|
||||||
|
{
|
||||||
|
chatWithPromptCalled: false,
|
||||||
|
exists: true,
|
||||||
|
title: null,
|
||||||
|
}
|
||||||
|
|
||||||
|
> should use correct prompt for title generation
|
||||||
|
|
||||||
|
{
|
||||||
|
content: `[user]: Explain quantum computing briefly␊
|
||||||
|
[assistant]: Quantum computing uses quantum mechanics principles.`,
|
||||||
|
promptName: 'Summary as title',
|
||||||
|
}
|
||||||
|
|||||||
Binary file not shown.
@@ -11,7 +11,11 @@ import { EventBus, JobQueue } from '../base';
|
|||||||
import { ConfigModule } from '../base/config';
|
import { ConfigModule } from '../base/config';
|
||||||
import { AuthService } from '../core/auth';
|
import { AuthService } from '../core/auth';
|
||||||
import { QuotaModule } from '../core/quota';
|
import { QuotaModule } from '../core/quota';
|
||||||
import { ContextCategories, WorkspaceModel } from '../models';
|
import {
|
||||||
|
ContextCategories,
|
||||||
|
CopilotSessionModel,
|
||||||
|
WorkspaceModel,
|
||||||
|
} from '../models';
|
||||||
import { CopilotModule } from '../plugins/copilot';
|
import { CopilotModule } from '../plugins/copilot';
|
||||||
import { CopilotContextService } from '../plugins/copilot/context';
|
import { CopilotContextService } from '../plugins/copilot/context';
|
||||||
import {
|
import {
|
||||||
@@ -57,12 +61,13 @@ import { MockCopilotProvider } from './mocks';
|
|||||||
import { createTestingModule, TestingModule } from './utils';
|
import { createTestingModule, TestingModule } from './utils';
|
||||||
import { WorkflowTestCases } from './utils/copilot';
|
import { WorkflowTestCases } from './utils/copilot';
|
||||||
|
|
||||||
const test = ava as TestFn<{
|
type Context = {
|
||||||
auth: AuthService;
|
auth: AuthService;
|
||||||
module: TestingModule;
|
module: TestingModule;
|
||||||
db: PrismaClient;
|
db: PrismaClient;
|
||||||
event: EventBus;
|
event: EventBus;
|
||||||
workspace: WorkspaceModel;
|
workspace: WorkspaceModel;
|
||||||
|
copilotSession: CopilotSessionModel;
|
||||||
context: CopilotContextService;
|
context: CopilotContextService;
|
||||||
prompt: PromptService;
|
prompt: PromptService;
|
||||||
transcript: CopilotTranscriptionService;
|
transcript: CopilotTranscriptionService;
|
||||||
@@ -78,7 +83,8 @@ const test = ava as TestFn<{
|
|||||||
html: CopilotCheckHtmlExecutor;
|
html: CopilotCheckHtmlExecutor;
|
||||||
json: CopilotCheckJsonExecutor;
|
json: CopilotCheckJsonExecutor;
|
||||||
};
|
};
|
||||||
}>;
|
};
|
||||||
|
const test = ava as TestFn<Context>;
|
||||||
let userId: string;
|
let userId: string;
|
||||||
|
|
||||||
test.before(async t => {
|
test.before(async t => {
|
||||||
@@ -119,6 +125,7 @@ test.before(async t => {
|
|||||||
const db = module.get(PrismaClient);
|
const db = module.get(PrismaClient);
|
||||||
const event = module.get(EventBus);
|
const event = module.get(EventBus);
|
||||||
const workspace = module.get(WorkspaceModel);
|
const workspace = module.get(WorkspaceModel);
|
||||||
|
const copilotSession = module.get(CopilotSessionModel);
|
||||||
const prompt = module.get(PromptService);
|
const prompt = module.get(PromptService);
|
||||||
const factory = module.get(CopilotProviderFactory);
|
const factory = module.get(CopilotProviderFactory);
|
||||||
|
|
||||||
@@ -136,6 +143,7 @@ test.before(async t => {
|
|||||||
t.context.db = db;
|
t.context.db = db;
|
||||||
t.context.event = event;
|
t.context.event = event;
|
||||||
t.context.workspace = workspace;
|
t.context.workspace = workspace;
|
||||||
|
t.context.copilotSession = copilotSession;
|
||||||
t.context.prompt = prompt;
|
t.context.prompt = prompt;
|
||||||
t.context.factory = factory;
|
t.context.factory = factory;
|
||||||
t.context.session = session;
|
t.context.session = session;
|
||||||
@@ -1752,3 +1760,168 @@ test('should be able to manage workspace embedding', async t => {
|
|||||||
t.is(ret2.length, 0, 'should not match workspace context');
|
t.is(ret2.length, 0, 'should not match workspace context');
|
||||||
}
|
}
|
||||||
});
|
});
|
||||||
|
|
||||||
|
test('should handle generateSessionTitle correctly under various conditions', async t => {
|
||||||
|
const { prompt, session, workspace, copilotSession } = t.context;
|
||||||
|
|
||||||
|
await prompt.set('test', 'model', [{ role: 'user', content: '{{content}}' }]);
|
||||||
|
const createSession = async (
|
||||||
|
options: {
|
||||||
|
userMessage?: string;
|
||||||
|
assistantMessage?: string;
|
||||||
|
existingTitle?: string;
|
||||||
|
} = {}
|
||||||
|
) => {
|
||||||
|
const ws = await workspace.create(userId);
|
||||||
|
const sessionId = await session.create({
|
||||||
|
docId: 'test-doc',
|
||||||
|
workspaceId: ws.id,
|
||||||
|
userId,
|
||||||
|
promptName: 'test',
|
||||||
|
pinned: false,
|
||||||
|
});
|
||||||
|
|
||||||
|
if (options.existingTitle) {
|
||||||
|
await copilotSession.update({
|
||||||
|
userId,
|
||||||
|
sessionId,
|
||||||
|
title: options.existingTitle,
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
const chatSession = await session.get(sessionId);
|
||||||
|
if (chatSession) {
|
||||||
|
if (options.userMessage) {
|
||||||
|
chatSession.push({
|
||||||
|
role: 'user',
|
||||||
|
content: options.userMessage,
|
||||||
|
createdAt: new Date(),
|
||||||
|
});
|
||||||
|
}
|
||||||
|
if (options.assistantMessage) {
|
||||||
|
chatSession.push({
|
||||||
|
role: 'assistant',
|
||||||
|
content: options.assistantMessage,
|
||||||
|
createdAt: new Date(),
|
||||||
|
});
|
||||||
|
}
|
||||||
|
await chatSession.save();
|
||||||
|
}
|
||||||
|
|
||||||
|
return sessionId;
|
||||||
|
};
|
||||||
|
|
||||||
|
const testCases = [
|
||||||
|
{
|
||||||
|
name: 'should generate title when conditions are met',
|
||||||
|
setup: () =>
|
||||||
|
createSession({
|
||||||
|
userMessage: 'What is machine learning?',
|
||||||
|
assistantMessage:
|
||||||
|
'Machine learning is a subset of artificial intelligence.',
|
||||||
|
}),
|
||||||
|
mockFn: () => 'What is Machine Learning?',
|
||||||
|
expectSnapshot: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: 'should not generate title when session already has title',
|
||||||
|
setup: () =>
|
||||||
|
createSession({
|
||||||
|
userMessage: 'Test message',
|
||||||
|
assistantMessage: 'Test response',
|
||||||
|
existingTitle: 'Existing Title',
|
||||||
|
}),
|
||||||
|
mockFn: () => 'New Title',
|
||||||
|
expectSnapshot: true,
|
||||||
|
expectNotCalled: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: 'should not generate title when no user messages exist',
|
||||||
|
setup: () =>
|
||||||
|
createSession({ assistantMessage: 'Hello! How can I help you?' }),
|
||||||
|
mockFn: () => 'New Title',
|
||||||
|
expectSnapshot: true,
|
||||||
|
expectNotCalled: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: 'should not generate title when no assistant messages exist',
|
||||||
|
setup: () => createSession({ userMessage: 'What is AI?' }),
|
||||||
|
mockFn: () => 'New Title',
|
||||||
|
expectSnapshot: true,
|
||||||
|
expectNotCalled: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: 'should handle errors gracefully',
|
||||||
|
setup: () =>
|
||||||
|
createSession({
|
||||||
|
userMessage: 'Test question',
|
||||||
|
assistantMessage: 'Test answer',
|
||||||
|
}),
|
||||||
|
mockFn: () => {
|
||||||
|
throw new Error('Mock error for testing');
|
||||||
|
},
|
||||||
|
expectError: 'Mock error for testing',
|
||||||
|
},
|
||||||
|
];
|
||||||
|
|
||||||
|
for (const testCase of testCases) {
|
||||||
|
const sessionId = await testCase.setup();
|
||||||
|
let chatWithPromptCalled = false;
|
||||||
|
|
||||||
|
const mockStub = Sinon.stub(session, 'chatWithPrompt').callsFake(
|
||||||
|
async () => {
|
||||||
|
chatWithPromptCalled = true;
|
||||||
|
return testCase.mockFn();
|
||||||
|
}
|
||||||
|
);
|
||||||
|
|
||||||
|
if (testCase.expectError) {
|
||||||
|
await t.throwsAsync(
|
||||||
|
() => session.generateSessionTitle({ sessionId }),
|
||||||
|
{ message: testCase.expectError },
|
||||||
|
testCase.name
|
||||||
|
);
|
||||||
|
} else {
|
||||||
|
await session.generateSessionTitle({ sessionId });
|
||||||
|
|
||||||
|
if (testCase.expectSnapshot) {
|
||||||
|
const sessionState = await session.getSession(sessionId);
|
||||||
|
t.snapshot(
|
||||||
|
{
|
||||||
|
chatWithPromptCalled: testCase.expectNotCalled
|
||||||
|
? chatWithPromptCalled
|
||||||
|
: undefined,
|
||||||
|
title: sessionState?.title,
|
||||||
|
exists: !!sessionState,
|
||||||
|
},
|
||||||
|
testCase.name
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
mockStub.restore();
|
||||||
|
}
|
||||||
|
|
||||||
|
{
|
||||||
|
const sessionId = await createSession({
|
||||||
|
userMessage: 'Explain quantum computing briefly',
|
||||||
|
assistantMessage: 'Quantum computing uses quantum mechanics principles.',
|
||||||
|
});
|
||||||
|
|
||||||
|
let capturedArgs: any[] = [];
|
||||||
|
Sinon.stub(session, 'chatWithPrompt').callsFake(async (...args) => {
|
||||||
|
capturedArgs = args;
|
||||||
|
return 'Quantum Computing Explained';
|
||||||
|
});
|
||||||
|
|
||||||
|
await session.generateSessionTitle({ sessionId });
|
||||||
|
|
||||||
|
t.snapshot(
|
||||||
|
{
|
||||||
|
promptName: capturedArgs[0],
|
||||||
|
content: capturedArgs[1]?.content,
|
||||||
|
},
|
||||||
|
'should use correct prompt for title generation'
|
||||||
|
);
|
||||||
|
}
|
||||||
|
});
|
||||||
|
|||||||
@@ -58,6 +58,7 @@ test.beforeEach(async t => {
|
|||||||
workspaceId: workspace.id,
|
workspaceId: workspace.id,
|
||||||
docId,
|
docId,
|
||||||
userId: user.id,
|
userId: user.id,
|
||||||
|
title: null,
|
||||||
promptName: 'prompt-name',
|
promptName: 'prompt-name',
|
||||||
promptAction: null,
|
promptAction: null,
|
||||||
});
|
});
|
||||||
|
|||||||
@@ -82,6 +82,7 @@ const createTestSession = async (
|
|||||||
workspaceId: workspace.id,
|
workspaceId: workspace.id,
|
||||||
docId: null,
|
docId: null,
|
||||||
pinned: false,
|
pinned: false,
|
||||||
|
title: null,
|
||||||
promptName: TEST_PROMPTS.NORMAL,
|
promptName: TEST_PROMPTS.NORMAL,
|
||||||
promptAction: null,
|
promptAction: null,
|
||||||
...overrides,
|
...overrides,
|
||||||
@@ -297,6 +298,7 @@ test('should pin and unpin sessions', async t => {
|
|||||||
promptName: 'test-prompt',
|
promptName: 'test-prompt',
|
||||||
promptAction: null,
|
promptAction: null,
|
||||||
pinned: true,
|
pinned: true,
|
||||||
|
title: null,
|
||||||
});
|
});
|
||||||
|
|
||||||
const firstSession = await copilotSession.get(firstSessionId);
|
const firstSession = await copilotSession.get(firstSessionId);
|
||||||
@@ -312,6 +314,7 @@ test('should pin and unpin sessions', async t => {
|
|||||||
promptName: 'test-prompt',
|
promptName: 'test-prompt',
|
||||||
promptAction: null,
|
promptAction: null,
|
||||||
pinned: true,
|
pinned: true,
|
||||||
|
title: null,
|
||||||
});
|
});
|
||||||
|
|
||||||
const sessionStatesAfterSecondPin = await getSessionStates(db, [
|
const sessionStatesAfterSecondPin = await getSessionStates(db, [
|
||||||
@@ -796,6 +799,7 @@ test('should handle fork and session attachment operations', async t => {
|
|||||||
workspaceId: workspace.id,
|
workspaceId: workspace.id,
|
||||||
docId: forkConfig.docId,
|
docId: forkConfig.docId,
|
||||||
pinned: forkConfig.pinned,
|
pinned: forkConfig.pinned,
|
||||||
|
title: null,
|
||||||
parentSessionId,
|
parentSessionId,
|
||||||
prompt: { name: TEST_PROMPTS.NORMAL, action: null, model: 'gpt-4.1' },
|
prompt: { name: TEST_PROMPTS.NORMAL, action: null, model: 'gpt-4.1' },
|
||||||
messages: [
|
messages: [
|
||||||
|
|||||||
@@ -50,6 +50,7 @@ type PureChatSession = {
|
|||||||
workspaceId: string;
|
workspaceId: string;
|
||||||
docId?: string | null;
|
docId?: string | null;
|
||||||
pinned?: boolean;
|
pinned?: boolean;
|
||||||
|
title: string | null;
|
||||||
messages?: ChatMessage[];
|
messages?: ChatMessage[];
|
||||||
// connect ids
|
// connect ids
|
||||||
userId: string;
|
userId: string;
|
||||||
@@ -82,7 +83,7 @@ type UpdateChatSessionMessage = ChatSessionBaseState & {
|
|||||||
};
|
};
|
||||||
|
|
||||||
export type UpdateChatSessionOptions = ChatSessionBaseState &
|
export type UpdateChatSessionOptions = ChatSessionBaseState &
|
||||||
Pick<Partial<ChatSession>, 'docId' | 'pinned' | 'promptName'>;
|
Pick<Partial<ChatSession>, 'docId' | 'pinned' | 'promptName' | 'title'>;
|
||||||
|
|
||||||
export type UpdateChatSession = ChatSessionBaseState & UpdateChatSessionOptions;
|
export type UpdateChatSession = ChatSessionBaseState & UpdateChatSessionOptions;
|
||||||
|
|
||||||
@@ -254,7 +255,7 @@ export class CopilotSessionModel extends BaseModel {
|
|||||||
return (await this.db.aiSession.findUnique({
|
return (await this.db.aiSession.findUnique({
|
||||||
where: { ...where, id: sessionId, deletedAt: null },
|
where: { ...where, id: sessionId, deletedAt: null },
|
||||||
select,
|
select,
|
||||||
})) as Prisma.AiSessionGetPayload<{ select: Select }>;
|
})) as Prisma.AiSessionGetPayload<{ select: Select }> | null;
|
||||||
}
|
}
|
||||||
|
|
||||||
@Transactional()
|
@Transactional()
|
||||||
@@ -266,6 +267,7 @@ export class CopilotSessionModel extends BaseModel {
|
|||||||
docId: true,
|
docId: true,
|
||||||
pinned: true,
|
pinned: true,
|
||||||
parentSessionId: true,
|
parentSessionId: true,
|
||||||
|
title: true,
|
||||||
messages: {
|
messages: {
|
||||||
select: {
|
select: {
|
||||||
id: true,
|
id: true,
|
||||||
@@ -331,6 +333,7 @@ export class CopilotSessionModel extends BaseModel {
|
|||||||
docId: true,
|
docId: true,
|
||||||
parentSessionId: true,
|
parentSessionId: true,
|
||||||
pinned: true,
|
pinned: true,
|
||||||
|
title: true,
|
||||||
promptName: true,
|
promptName: true,
|
||||||
tokenCost: true,
|
tokenCost: true,
|
||||||
createdAt: true,
|
createdAt: true,
|
||||||
@@ -373,7 +376,7 @@ export class CopilotSessionModel extends BaseModel {
|
|||||||
|
|
||||||
@Transactional()
|
@Transactional()
|
||||||
async update(options: UpdateChatSessionOptions): Promise<string> {
|
async update(options: UpdateChatSessionOptions): Promise<string> {
|
||||||
const { userId, sessionId, docId, promptName, pinned } = options;
|
const { userId, sessionId, docId, promptName, pinned, title } = options;
|
||||||
const session = await this.getExists(
|
const session = await this.getExists(
|
||||||
sessionId,
|
sessionId,
|
||||||
{
|
{
|
||||||
@@ -419,7 +422,7 @@ export class CopilotSessionModel extends BaseModel {
|
|||||||
|
|
||||||
await this.db.aiSession.update({
|
await this.db.aiSession.update({
|
||||||
where: { id: sessionId },
|
where: { id: sessionId },
|
||||||
data: { docId, promptName, pinned },
|
data: { docId, promptName, pinned, title },
|
||||||
});
|
});
|
||||||
|
|
||||||
return sessionId;
|
return sessionId;
|
||||||
@@ -522,17 +525,29 @@ export class CopilotSessionModel extends BaseModel {
|
|||||||
if (!id) {
|
if (!id) {
|
||||||
throw new CopilotSessionNotFound();
|
throw new CopilotSessionNotFound();
|
||||||
}
|
}
|
||||||
const ids = await this.getMessages(id, { id: true, role: true }).then(
|
const messages = await this.getMessages(id, { id: true, role: true });
|
||||||
roles =>
|
const ids = messages
|
||||||
roles
|
|
||||||
.slice(
|
.slice(
|
||||||
roles.findLastIndex(({ role }) => role === AiPromptRole.user) +
|
messages.findLastIndex(({ role }) => role === AiPromptRole.user) +
|
||||||
(removeLatestUserMessage ? 0 : 1)
|
(removeLatestUserMessage ? 0 : 1)
|
||||||
)
|
)
|
||||||
.map(({ id }) => id)
|
.map(({ id }) => id);
|
||||||
);
|
|
||||||
if (ids.length) {
|
if (ids.length) {
|
||||||
await this.db.aiSessionMessage.deleteMany({ where: { id: { in: ids } } });
|
await this.db.aiSessionMessage.deleteMany({ where: { id: { in: ids } } });
|
||||||
|
|
||||||
|
// clear the title if there only one round of conversation left
|
||||||
|
const remainingMessages = await this.getMessages(id, { role: true });
|
||||||
|
const userMessageCount = remainingMessages.filter(
|
||||||
|
m => m.role === AiPromptRole.user
|
||||||
|
).length;
|
||||||
|
|
||||||
|
if (userMessageCount <= 1) {
|
||||||
|
await this.db.aiSession.update({
|
||||||
|
where: { id },
|
||||||
|
data: { title: null },
|
||||||
|
});
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -67,7 +67,9 @@ class CreateChatSessionInput {
|
|||||||
}
|
}
|
||||||
|
|
||||||
@InputType()
|
@InputType()
|
||||||
class UpdateChatSessionInput implements Omit<UpdateChatSession, 'userId'> {
|
class UpdateChatSessionInput
|
||||||
|
implements Omit<UpdateChatSession, 'userId' | 'title'>
|
||||||
|
{
|
||||||
@Field(() => String)
|
@Field(() => String)
|
||||||
sessionId!: string;
|
sessionId!: string;
|
||||||
|
|
||||||
@@ -336,6 +338,9 @@ export class CopilotSessionType {
|
|||||||
@Field(() => Boolean)
|
@Field(() => Boolean)
|
||||||
pinned!: boolean;
|
pinned!: boolean;
|
||||||
|
|
||||||
|
@Field(() => String, { nullable: true })
|
||||||
|
title!: string | null;
|
||||||
|
|
||||||
@Field(() => ID, { nullable: true })
|
@Field(() => ID, { nullable: true })
|
||||||
parentSessionId!: string | null;
|
parentSessionId!: string | null;
|
||||||
|
|
||||||
@@ -653,6 +658,7 @@ export class CopilotResolver {
|
|||||||
parentSessionId: session.parentSessionId,
|
parentSessionId: session.parentSessionId,
|
||||||
docId: session.docId,
|
docId: session.docId,
|
||||||
pinned: session.pinned,
|
pinned: session.pinned,
|
||||||
|
title: session.title,
|
||||||
promptName: session.prompt.name,
|
promptName: session.prompt.name,
|
||||||
model: session.prompt.model,
|
model: session.prompt.model,
|
||||||
optionalModels: session.prompt.optionalModels,
|
optionalModels: session.prompt.optionalModels,
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
import { randomUUID } from 'node:crypto';
|
import { randomUUID } from 'node:crypto';
|
||||||
|
|
||||||
import { Injectable, Logger } from '@nestjs/common';
|
import { Injectable, Logger } from '@nestjs/common';
|
||||||
|
import { ModuleRef } from '@nestjs/core';
|
||||||
import { Transactional } from '@nestjs-cls/transactional';
|
import { Transactional } from '@nestjs-cls/transactional';
|
||||||
import { AiPromptRole } from '@prisma/client';
|
import { AiPromptRole } from '@prisma/client';
|
||||||
|
|
||||||
@@ -11,6 +12,9 @@ import {
|
|||||||
CopilotQuotaExceeded,
|
CopilotQuotaExceeded,
|
||||||
CopilotSessionInvalidInput,
|
CopilotSessionInvalidInput,
|
||||||
CopilotSessionNotFound,
|
CopilotSessionNotFound,
|
||||||
|
JobQueue,
|
||||||
|
NoCopilotProviderAvailable,
|
||||||
|
OnJob,
|
||||||
} from '../../base';
|
} from '../../base';
|
||||||
import { QuotaService } from '../../core/quota';
|
import { QuotaService } from '../../core/quota';
|
||||||
import {
|
import {
|
||||||
@@ -22,7 +26,12 @@ import {
|
|||||||
} from '../../models';
|
} from '../../models';
|
||||||
import { ChatMessageCache } from './message';
|
import { ChatMessageCache } from './message';
|
||||||
import { PromptService } from './prompt';
|
import { PromptService } from './prompt';
|
||||||
import { PromptMessage, PromptParams } from './providers';
|
import {
|
||||||
|
CopilotProviderFactory,
|
||||||
|
ModelOutputType,
|
||||||
|
PromptMessage,
|
||||||
|
PromptParams,
|
||||||
|
} from './providers';
|
||||||
import {
|
import {
|
||||||
type ChatHistory,
|
type ChatHistory,
|
||||||
type ChatMessage,
|
type ChatMessage,
|
||||||
@@ -33,6 +42,14 @@ import {
|
|||||||
type SubmittedMessage,
|
type SubmittedMessage,
|
||||||
} from './types';
|
} from './types';
|
||||||
|
|
||||||
|
declare global {
|
||||||
|
interface Jobs {
|
||||||
|
'copilot.session.generateTitle': {
|
||||||
|
sessionId: string;
|
||||||
|
};
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
export class ChatSession implements AsyncDisposable {
|
export class ChatSession implements AsyncDisposable {
|
||||||
private stashMessageCount = 0;
|
private stashMessageCount = 0;
|
||||||
constructor(
|
constructor(
|
||||||
@@ -224,10 +241,12 @@ export class ChatSessionService {
|
|||||||
private readonly logger = new Logger(ChatSessionService.name);
|
private readonly logger = new Logger(ChatSessionService.name);
|
||||||
|
|
||||||
constructor(
|
constructor(
|
||||||
|
private readonly moduleRef: ModuleRef,
|
||||||
|
private readonly models: Models,
|
||||||
|
private readonly jobs: JobQueue,
|
||||||
private readonly quota: QuotaService,
|
private readonly quota: QuotaService,
|
||||||
private readonly messageCache: ChatMessageCache,
|
private readonly messageCache: ChatMessageCache,
|
||||||
private readonly prompt: PromptService,
|
private readonly prompt: PromptService
|
||||||
private readonly models: Models
|
|
||||||
) {}
|
) {}
|
||||||
|
|
||||||
async getSession(sessionId: string): Promise<ChatSessionState | undefined> {
|
async getSession(sessionId: string): Promise<ChatSessionState | undefined> {
|
||||||
@@ -244,6 +263,7 @@ export class ChatSessionService {
|
|||||||
workspaceId: session.workspaceId,
|
workspaceId: session.workspaceId,
|
||||||
docId: session.docId,
|
docId: session.docId,
|
||||||
pinned: session.pinned,
|
pinned: session.pinned,
|
||||||
|
title: session.title,
|
||||||
parentSessionId: session.parentSessionId,
|
parentSessionId: session.parentSessionId,
|
||||||
prompt,
|
prompt,
|
||||||
messages: messages.success ? messages.data : [],
|
messages: messages.success ? messages.data : [],
|
||||||
@@ -282,6 +302,7 @@ export class ChatSessionService {
|
|||||||
workspaceId: session.workspaceId,
|
workspaceId: session.workspaceId,
|
||||||
docId: session.docId,
|
docId: session.docId,
|
||||||
pinned: session.pinned,
|
pinned: session.pinned,
|
||||||
|
title: session.title,
|
||||||
parentSessionId: session.parentSessionId,
|
parentSessionId: session.parentSessionId,
|
||||||
prompt,
|
prompt,
|
||||||
};
|
};
|
||||||
@@ -303,6 +324,7 @@ export class ChatSessionService {
|
|||||||
workspaceId,
|
workspaceId,
|
||||||
docId,
|
docId,
|
||||||
pinned,
|
pinned,
|
||||||
|
title,
|
||||||
promptName,
|
promptName,
|
||||||
tokenCost,
|
tokenCost,
|
||||||
messages,
|
messages,
|
||||||
@@ -347,6 +369,7 @@ export class ChatSessionService {
|
|||||||
workspaceId,
|
workspaceId,
|
||||||
docId,
|
docId,
|
||||||
pinned,
|
pinned,
|
||||||
|
title,
|
||||||
action: prompt.action || null,
|
action: prompt.action || null,
|
||||||
tokens: tokenCost,
|
tokens: tokenCost,
|
||||||
createdAt,
|
createdAt,
|
||||||
@@ -418,6 +441,7 @@ export class ChatSessionService {
|
|||||||
...options,
|
...options,
|
||||||
sessionId,
|
sessionId,
|
||||||
prompt,
|
prompt,
|
||||||
|
title: null,
|
||||||
messages: [],
|
messages: [],
|
||||||
// when client create chat session, we always find root session
|
// when client create chat session, we always find root session
|
||||||
parentSessionId: null,
|
parentSessionId: null,
|
||||||
@@ -520,8 +544,78 @@ export class ChatSessionService {
|
|||||||
if (state) {
|
if (state) {
|
||||||
return new ChatSession(this.messageCache, state, async state => {
|
return new ChatSession(this.messageCache, state, async state => {
|
||||||
await this.models.copilotSession.updateMessages(state);
|
await this.models.copilotSession.updateMessages(state);
|
||||||
|
if (!state.prompt.action) {
|
||||||
|
await this.jobs.add('copilot.session.generateTitle', { sessionId });
|
||||||
|
}
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
return null;
|
return null;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// public for test mock
|
||||||
|
async chatWithPrompt(
|
||||||
|
promptName: string,
|
||||||
|
message: Partial<PromptMessage>
|
||||||
|
): Promise<string> {
|
||||||
|
const prompt = await this.prompt.get(promptName);
|
||||||
|
if (!prompt) {
|
||||||
|
throw new CopilotPromptNotFound({ name: promptName });
|
||||||
|
}
|
||||||
|
|
||||||
|
const cond = { modelId: prompt.model };
|
||||||
|
const msg = { role: 'user' as const, content: '', ...message };
|
||||||
|
const config = Object.assign({}, prompt.config);
|
||||||
|
|
||||||
|
const provider = await this.moduleRef
|
||||||
|
.get(CopilotProviderFactory)
|
||||||
|
.getProvider({
|
||||||
|
outputType: ModelOutputType.Text,
|
||||||
|
modelId: prompt.model,
|
||||||
|
});
|
||||||
|
|
||||||
|
if (!provider) {
|
||||||
|
throw new NoCopilotProviderAvailable();
|
||||||
|
}
|
||||||
|
|
||||||
|
return provider.text(cond, [...prompt.finish({}), msg], config);
|
||||||
|
}
|
||||||
|
|
||||||
|
@OnJob('copilot.session.generateTitle')
|
||||||
|
async generateSessionTitle(job: Jobs['copilot.session.generateTitle']) {
|
||||||
|
const { sessionId } = job;
|
||||||
|
|
||||||
|
try {
|
||||||
|
const session = await this.models.copilotSession.get(sessionId);
|
||||||
|
if (!session) {
|
||||||
|
this.logger.warn(
|
||||||
|
`Session ${sessionId} not found when generating title`
|
||||||
|
);
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
const { userId, title, messages } = session;
|
||||||
|
if (
|
||||||
|
title ||
|
||||||
|
!messages.length ||
|
||||||
|
messages.filter(m => m.role === 'user').length === 0 ||
|
||||||
|
messages.filter(m => m.role === 'assistant').length === 0
|
||||||
|
) {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
{
|
||||||
|
const title = await this.chatWithPrompt('Summary as title', {
|
||||||
|
content: session.messages
|
||||||
|
.map(m => `[${m.role}]: ${m.content}`)
|
||||||
|
.join('\n'),
|
||||||
|
});
|
||||||
|
await this.models.copilotSession.update({ userId, sessionId, title });
|
||||||
|
}
|
||||||
|
} catch (error) {
|
||||||
|
console.error(
|
||||||
|
`Failed to generate title for session ${sessionId}:`,
|
||||||
|
error
|
||||||
|
);
|
||||||
|
throw error;
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -50,6 +50,7 @@ export const ChatHistorySchema = z
|
|||||||
workspaceId: z.string(),
|
workspaceId: z.string(),
|
||||||
docId: z.string().nullable(),
|
docId: z.string().nullable(),
|
||||||
pinned: z.boolean(),
|
pinned: z.boolean(),
|
||||||
|
title: z.string().nullable(),
|
||||||
action: z.string().nullable(),
|
action: z.string().nullable(),
|
||||||
tokens: z.number(),
|
tokens: z.number(),
|
||||||
messages: z.array(ChatMessageSchema),
|
messages: z.array(ChatMessageSchema),
|
||||||
@@ -85,6 +86,7 @@ export interface ChatSessionForkOptions
|
|||||||
|
|
||||||
export interface ChatSessionState
|
export interface ChatSessionState
|
||||||
extends Omit<ChatSessionOptions, 'promptName'> {
|
extends Omit<ChatSessionOptions, 'promptName'> {
|
||||||
|
title: string | null;
|
||||||
// connect ids
|
// connect ids
|
||||||
sessionId: string;
|
sessionId: string;
|
||||||
parentSessionId: string | null;
|
parentSessionId: string | null;
|
||||||
|
|||||||
@@ -324,6 +324,7 @@ type CopilotSessionType {
|
|||||||
parentSessionId: ID
|
parentSessionId: ID
|
||||||
pinned: Boolean!
|
pinned: Boolean!
|
||||||
promptName: String!
|
promptName: String!
|
||||||
|
title: String
|
||||||
}
|
}
|
||||||
|
|
||||||
type CopilotWorkspaceConfig {
|
type CopilotWorkspaceConfig {
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ query getCopilotSession(
|
|||||||
parentSessionId
|
parentSessionId
|
||||||
docId
|
docId
|
||||||
pinned
|
pinned
|
||||||
|
title
|
||||||
promptName
|
promptName
|
||||||
model
|
model
|
||||||
optionalModels
|
optionalModels
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ query getCopilotSessions(
|
|||||||
parentSessionId
|
parentSessionId
|
||||||
docId
|
docId
|
||||||
pinned
|
pinned
|
||||||
|
title
|
||||||
promptName
|
promptName
|
||||||
model
|
model
|
||||||
optionalModels
|
optionalModels
|
||||||
|
|||||||
@@ -799,6 +799,7 @@ export const getCopilotSessionQuery = {
|
|||||||
parentSessionId
|
parentSessionId
|
||||||
docId
|
docId
|
||||||
pinned
|
pinned
|
||||||
|
title
|
||||||
promptName
|
promptName
|
||||||
model
|
model
|
||||||
optionalModels
|
optionalModels
|
||||||
@@ -848,6 +849,7 @@ export const getCopilotSessionsQuery = {
|
|||||||
parentSessionId
|
parentSessionId
|
||||||
docId
|
docId
|
||||||
pinned
|
pinned
|
||||||
|
title
|
||||||
promptName
|
promptName
|
||||||
model
|
model
|
||||||
optionalModels
|
optionalModels
|
||||||
|
|||||||
@@ -419,6 +419,7 @@ export interface CopilotSessionType {
|
|||||||
parentSessionId: Maybe<Scalars['ID']['output']>;
|
parentSessionId: Maybe<Scalars['ID']['output']>;
|
||||||
pinned: Scalars['Boolean']['output'];
|
pinned: Scalars['Boolean']['output'];
|
||||||
promptName: Scalars['String']['output'];
|
promptName: Scalars['String']['output'];
|
||||||
|
title: Maybe<Scalars['String']['output']>;
|
||||||
}
|
}
|
||||||
|
|
||||||
export interface CopilotWorkspaceConfig {
|
export interface CopilotWorkspaceConfig {
|
||||||
@@ -3619,6 +3620,7 @@ export type GetCopilotSessionQuery = {
|
|||||||
parentSessionId: string | null;
|
parentSessionId: string | null;
|
||||||
docId: string | null;
|
docId: string | null;
|
||||||
pinned: boolean;
|
pinned: boolean;
|
||||||
|
title: string | null;
|
||||||
promptName: string;
|
promptName: string;
|
||||||
model: string;
|
model: string;
|
||||||
optionalModels: Array<string>;
|
optionalModels: Array<string>;
|
||||||
@@ -3680,6 +3682,7 @@ export type GetCopilotSessionsQuery = {
|
|||||||
parentSessionId: string | null;
|
parentSessionId: string | null;
|
||||||
docId: string | null;
|
docId: string | null;
|
||||||
pinned: boolean;
|
pinned: boolean;
|
||||||
|
title: string | null;
|
||||||
promptName: string;
|
promptName: string;
|
||||||
model: string;
|
model: string;
|
||||||
optionalModels: Array<string>;
|
optionalModels: Array<string>;
|
||||||
|
|||||||
Reference in New Issue
Block a user