496 lines
17 KiB
TypeScript
496 lines
17 KiB
TypeScript
import { EntityManager } from "@mikro-orm/postgresql";
|
|
import { BadRequestException, Injectable, Logger, NotFoundException } from "@nestjs/common";
|
|
import { ConfigService } from "@nestjs/config";
|
|
|
|
import { LangChainService } from "./langchain.service";
|
|
import { LLMService } from "./llm.service";
|
|
import { User } from "../../users/entities/user.entity";
|
|
// import { UsersService } from "../../users/services/users.service";
|
|
import { CreateChatSessionDto } from "../DTO/create-chat-session.dto";
|
|
import { SendMessageDto } from "../DTO/send-message.dto";
|
|
import { ChatMessage, MessageStatus, MessageType } from "../entities/chat-message.entity";
|
|
import { ChatSession, ChatSessionStatus } from "../entities/chat-session.entity";
|
|
import { IChatContext } from "../interfaces/chatbot.interface";
|
|
import { ChatMessageRepository } from "../repositories/chat-message.repository";
|
|
import { ChatSessionRepository } from "../repositories/chat-session.repository";
|
|
|
|
export enum LLMProvider {
|
|
LANGCHAIN = "langchain",
|
|
GEMINI = "gemini",
|
|
}
|
|
|
|
@Injectable()
|
|
export class ChatbotService {
|
|
private readonly logger = new Logger(ChatbotService.name);
|
|
private defaultProvider: LLMProvider;
|
|
|
|
constructor(
|
|
private em: EntityManager,
|
|
private llmService: LLMService,
|
|
private langChainService: LangChainService,
|
|
private chatSessionRepo: ChatSessionRepository,
|
|
private chatMessageRepo: ChatMessageRepository,
|
|
private configService: ConfigService,
|
|
// private usersService: UsersService,
|
|
) {
|
|
// Configure default provider - LangChain for enhanced responses
|
|
this.defaultProvider = this.configService.get("DEFAULT_LLM_PROVIDER") === "gemini" ? LLMProvider.GEMINI : LLMProvider.LANGCHAIN;
|
|
|
|
this.logger.log(`Using ${this.defaultProvider} as default LLM provider`);
|
|
}
|
|
|
|
async createChatSession(userId: string, createDto: CreateChatSessionDto) {
|
|
const em = this.em.fork();
|
|
const user = await em.findOne(User, { id: userId });
|
|
if (!user) {
|
|
throw new NotFoundException("User not found");
|
|
}
|
|
|
|
const session = new ChatSession();
|
|
session.title = createDto.title;
|
|
session.user = user;
|
|
session.context = createDto.context || {};
|
|
session.lastMessageAt = new Date();
|
|
|
|
await em.persistAndFlush(session);
|
|
|
|
return this.mapSessionToDto(session);
|
|
}
|
|
//************************************ */
|
|
async getUserChatSessions(userId: string, limit = 10) {
|
|
const em = this.em.fork();
|
|
const sessions = await this.chatSessionRepo.findByUserWithMessages(userId, limit, em);
|
|
return sessions.map((session) => this.mapSessionToDto(session));
|
|
}
|
|
//************************************ */
|
|
|
|
async getChatSession(sessionId: string, userId: string) {
|
|
const em = this.em.fork();
|
|
const session = await em.findOne(ChatSession, { id: sessionId, user: userId }, { populate: ["messages", "messages.sender"] });
|
|
|
|
if (!session) {
|
|
throw new NotFoundException("Chat session not found");
|
|
}
|
|
|
|
return this.mapSessionToDto(session);
|
|
}
|
|
//************************************ */
|
|
|
|
async sendMessage(userId: string, sendDto: SendMessageDto, provider?: LLMProvider) {
|
|
const selectedProvider = provider || this.defaultProvider;
|
|
|
|
// Use LangChain by default for enhanced responses
|
|
if (selectedProvider === LLMProvider.LANGCHAIN) {
|
|
try {
|
|
return await this.sendMessageWithLangChain(userId, sendDto);
|
|
} catch (error) {
|
|
this.logger.error("LangChain failed, falling back to Gemini", error);
|
|
// Fallback to Gemini service
|
|
return await this.sendMessageWithGemini(userId, sendDto);
|
|
}
|
|
} else {
|
|
return await this.sendMessageWithGemini(userId, sendDto);
|
|
}
|
|
}
|
|
|
|
//************************************ */
|
|
|
|
private async sendMessageWithGemini(userId: string, sendDto: SendMessageDto) {
|
|
const em = this.em.fork();
|
|
const session = await em.findOne(ChatSession, { id: sendDto.sessionId, user: userId }, { populate: ["messages"] });
|
|
|
|
if (!session) {
|
|
throw new NotFoundException("Chat session not found");
|
|
}
|
|
|
|
if (session.status !== ChatSessionStatus.ACTIVE) {
|
|
throw new BadRequestException("Cannot send message to inactive session");
|
|
}
|
|
|
|
const user = await em.findOne(User, { id: userId });
|
|
if (!user) {
|
|
throw new NotFoundException("User not found");
|
|
}
|
|
|
|
// Create user message
|
|
const userMessage = new ChatMessage();
|
|
userMessage.content = sendDto.content;
|
|
userMessage.type = MessageType.USER;
|
|
userMessage.session = session;
|
|
userMessage.sender = user;
|
|
userMessage.responseToId = sendDto.responseToId;
|
|
userMessage.metadata = sendDto.metadata;
|
|
|
|
await em.persistAndFlush(userMessage);
|
|
|
|
// Update session last message time
|
|
await this.chatSessionRepo.updateLastMessageTime(session.id, em);
|
|
|
|
// Generate bot response asynchronously using original Gemini service
|
|
this.generateBotResponse(session, userMessage, userId);
|
|
|
|
return this.mapMessageToDto(userMessage);
|
|
}
|
|
|
|
//************************************ */
|
|
|
|
async sendMessageStream(userId: string, sendDto: SendMessageDto, provider?: LLMProvider) {
|
|
const selectedProvider = provider || this.defaultProvider;
|
|
|
|
// Use LangChain by default for enhanced streaming responses
|
|
if (selectedProvider === LLMProvider.LANGCHAIN) {
|
|
try {
|
|
return await this.sendMessageStreamWithLangChain(userId, sendDto);
|
|
} catch (error) {
|
|
this.logger.error("LangChain streaming failed, falling back to Gemini", error);
|
|
// Fallback to Gemini service
|
|
return await this.sendMessageStreamWithGemini(userId, sendDto);
|
|
}
|
|
} else {
|
|
return await this.sendMessageStreamWithGemini(userId, sendDto);
|
|
}
|
|
}
|
|
|
|
//************************************ */
|
|
|
|
private async sendMessageStreamWithGemini(userId: string, sendDto: SendMessageDto) {
|
|
const em = this.em.fork();
|
|
const session = await em.findOne(ChatSession, { id: sendDto.sessionId, user: userId }, { populate: ["messages"] });
|
|
|
|
if (!session) {
|
|
throw new NotFoundException("Chat session not found");
|
|
}
|
|
|
|
if (session.status !== ChatSessionStatus.ACTIVE) {
|
|
throw new BadRequestException("Cannot send message to inactive session");
|
|
}
|
|
|
|
const user = await em.findOne(User, { id: userId });
|
|
if (!user) {
|
|
throw new NotFoundException("User not found");
|
|
}
|
|
|
|
// Create user message
|
|
const userMessage = new ChatMessage();
|
|
userMessage.content = sendDto.content;
|
|
userMessage.type = MessageType.USER;
|
|
userMessage.session = session;
|
|
userMessage.sender = user;
|
|
userMessage.responseToId = sendDto.responseToId;
|
|
userMessage.metadata = sendDto.metadata;
|
|
|
|
await em.persistAndFlush(userMessage);
|
|
|
|
// Update session last message time
|
|
await this.chatSessionRepo.updateLastMessageTime(session.id, em);
|
|
|
|
// Get conversation history for context
|
|
const history = await this.chatMessageRepo.getConversationHistory(session.id, 20, em);
|
|
|
|
// Build context for LLM
|
|
const context: IChatContext = {
|
|
userId,
|
|
sessionId: session.id,
|
|
conversationHistory: history.reverse().map((msg) => ({
|
|
content: msg.content,
|
|
type: msg.type as "user" | "bot" | "system",
|
|
timestamp: msg.createdAt,
|
|
metadata: msg.metadata,
|
|
})),
|
|
userPreferences: session.context,
|
|
};
|
|
|
|
// Return a function that generates the stream using original LLM service
|
|
const streamGenerator = () => this.llmService.generateStreamResponse(sendDto.content, context);
|
|
|
|
return {
|
|
userMessage: this.mapMessageToDto(userMessage),
|
|
streamGenerator,
|
|
};
|
|
}
|
|
|
|
//************************************ */
|
|
|
|
private async generateBotResponse(session: ChatSession, userMessage: ChatMessage, userId: string): Promise<void> {
|
|
const em = this.em.fork();
|
|
try {
|
|
// Get conversation history
|
|
const history = await this.chatMessageRepo.getConversationHistory(session.id, 20, em);
|
|
|
|
// Build context for LLM
|
|
const context: IChatContext = {
|
|
userId,
|
|
sessionId: session.id,
|
|
conversationHistory: history.reverse().map((msg) => ({
|
|
content: msg.content,
|
|
type: msg.type as "user" | "bot" | "system",
|
|
timestamp: msg.createdAt,
|
|
metadata: msg.metadata,
|
|
})),
|
|
userPreferences: session.context,
|
|
};
|
|
|
|
// Generate response using LLM
|
|
const llmResponse = await this.llmService.generateResponse(userMessage.content, context);
|
|
|
|
// Create bot message
|
|
const botMessage = new ChatMessage();
|
|
botMessage.content = llmResponse.message;
|
|
botMessage.type = MessageType.BOT;
|
|
botMessage.session = session;
|
|
botMessage.responseToId = userMessage.id;
|
|
botMessage.tokensUsed = llmResponse.tokensUsed;
|
|
botMessage.metadata = {
|
|
confidence: llmResponse.confidence,
|
|
sources: llmResponse.sources,
|
|
llmContext: llmResponse.context,
|
|
};
|
|
|
|
await em.persistAndFlush(botMessage);
|
|
|
|
// Update session last message time
|
|
await this.chatSessionRepo.updateLastMessageTime(session.id, em);
|
|
|
|
this.logger.log(`Generated bot response for session ${session.id}`);
|
|
} catch (error) {
|
|
this.logger.error(`Failed to generate bot response for session ${session.id}`, error);
|
|
|
|
// Create error message
|
|
const errorMessage = new ChatMessage();
|
|
errorMessage.content = "متأسفم، در حال حاضر مشکلی در پاسخگویی دارم. لطفاً دوباره تلاش کنید. 🙏";
|
|
errorMessage.type = MessageType.BOT;
|
|
errorMessage.session = session;
|
|
errorMessage.responseToId = userMessage.id;
|
|
errorMessage.status = MessageStatus.FAILED;
|
|
|
|
await em.persistAndFlush(errorMessage);
|
|
}
|
|
}
|
|
|
|
//************************************ */
|
|
|
|
async closeChatSession(sessionId: string, userId: string) {
|
|
const em = this.em.fork();
|
|
const session = await em.findOne(ChatSession, { id: sessionId, user: userId });
|
|
|
|
if (!session) {
|
|
throw new NotFoundException("Chat session not found");
|
|
}
|
|
|
|
session.status = ChatSessionStatus.CLOSED;
|
|
await em.persistAndFlush(session);
|
|
|
|
return { message: "Chat session closed successfully" };
|
|
}
|
|
|
|
//************************************ */
|
|
|
|
async markMessagesAsRead(sessionId: string, userId: string, messageIds: string[]) {
|
|
const em = this.em.fork();
|
|
const session = await em.findOne(ChatSession, { id: sessionId, user: userId });
|
|
|
|
if (!session) {
|
|
throw new NotFoundException("Chat session not found");
|
|
}
|
|
|
|
await this.chatMessageRepo.markAsRead(messageIds, em);
|
|
|
|
return { message: "Messages marked as read successfully" };
|
|
}
|
|
|
|
//************************************ */
|
|
|
|
private mapSessionToDto(session: ChatSession) {
|
|
return {
|
|
id: session.id,
|
|
title: session.title,
|
|
status: session.status,
|
|
createdAt: session.createdAt,
|
|
lastMessageAt: session.lastMessageAt,
|
|
context: session.context,
|
|
messages: session.messages ? session.messages.getItems().map((msg) => this.mapMessageToDto(msg)) : [],
|
|
};
|
|
}
|
|
|
|
//************************************ */
|
|
|
|
private mapMessageToDto(message: ChatMessage) {
|
|
return {
|
|
id: message.id,
|
|
content: message.content,
|
|
type: message.type,
|
|
status: message.status,
|
|
createdAt: message.createdAt,
|
|
responseToId: message.responseToId,
|
|
metadata: message.metadata,
|
|
tokensUsed: message.tokensUsed,
|
|
};
|
|
}
|
|
|
|
//************************************ */
|
|
|
|
async sendMessageWithLangChain(userId: string, sendDto: SendMessageDto) {
|
|
const em = this.em.fork();
|
|
const session = await em.findOne(ChatSession, { id: sendDto.sessionId, user: userId }, { populate: ["messages"] });
|
|
|
|
if (!session) {
|
|
throw new NotFoundException("Chat session not found");
|
|
}
|
|
|
|
if (session.status !== ChatSessionStatus.ACTIVE) {
|
|
throw new BadRequestException("Cannot send message to inactive session");
|
|
}
|
|
|
|
const user = await em.findOne(User, { id: userId });
|
|
if (!user) {
|
|
throw new NotFoundException("User not found");
|
|
}
|
|
|
|
// Create user message
|
|
const userMessage = new ChatMessage();
|
|
userMessage.content = sendDto.content;
|
|
userMessage.type = MessageType.USER;
|
|
userMessage.session = session;
|
|
userMessage.sender = user;
|
|
userMessage.responseToId = sendDto.responseToId;
|
|
userMessage.metadata = sendDto.metadata;
|
|
|
|
await em.persistAndFlush(userMessage);
|
|
|
|
// Update session last message time
|
|
await this.chatSessionRepo.updateLastMessageTime(session.id, em);
|
|
|
|
// Generate bot response using LangChain
|
|
this.generateBotResponseWithLangChain(session, userMessage, userId);
|
|
|
|
return this.mapMessageToDto(userMessage);
|
|
}
|
|
|
|
//************************************ */
|
|
|
|
async sendMessageStreamWithLangChain(userId: string, sendDto: SendMessageDto) {
|
|
const em = this.em.fork();
|
|
const session = await em.findOne(ChatSession, { id: sendDto.sessionId, user: userId }, { populate: ["messages"] });
|
|
|
|
if (!session) {
|
|
throw new NotFoundException("Chat session not found");
|
|
}
|
|
|
|
if (session.status !== ChatSessionStatus.ACTIVE) {
|
|
throw new BadRequestException("Cannot send message to inactive session");
|
|
}
|
|
|
|
const user = await em.findOne(User, { id: userId });
|
|
if (!user) {
|
|
throw new NotFoundException("User not found");
|
|
}
|
|
|
|
// Create user message
|
|
const userMessage = new ChatMessage();
|
|
userMessage.content = sendDto.content;
|
|
userMessage.type = MessageType.USER;
|
|
userMessage.session = session;
|
|
userMessage.sender = user;
|
|
userMessage.responseToId = sendDto.responseToId;
|
|
userMessage.metadata = sendDto.metadata;
|
|
|
|
await em.persistAndFlush(userMessage);
|
|
|
|
// Update session last message time
|
|
await this.chatSessionRepo.updateLastMessageTime(session.id, em);
|
|
|
|
// Get conversation history for context
|
|
const history = await this.chatMessageRepo.getConversationHistory(session.id, 20, em);
|
|
|
|
// Build context for LangChain
|
|
const context: IChatContext = {
|
|
userId,
|
|
sessionId: session.id,
|
|
conversationHistory: history.reverse().map((msg) => ({
|
|
content: msg.content,
|
|
type: msg.type as "user" | "bot" | "system",
|
|
timestamp: msg.createdAt,
|
|
metadata: msg.metadata,
|
|
})),
|
|
userPreferences: session.context,
|
|
};
|
|
|
|
// Return a function that generates the stream using LangChain
|
|
const streamGenerator = () => this.langChainService.generateStreamResponse(sendDto.content, context);
|
|
|
|
return {
|
|
userMessage: this.mapMessageToDto(userMessage),
|
|
streamGenerator,
|
|
};
|
|
}
|
|
|
|
//************************************ */
|
|
|
|
private async generateBotResponseWithLangChain(session: ChatSession, userMessage: ChatMessage, userId: string): Promise<void> {
|
|
const em = this.em.fork();
|
|
try {
|
|
// Get conversation history
|
|
const history = await this.chatMessageRepo.getConversationHistory(session.id, 20, em);
|
|
|
|
// Build context for LangChain
|
|
const context: IChatContext = {
|
|
userId,
|
|
sessionId: session.id,
|
|
conversationHistory: history.reverse().map((msg) => ({
|
|
content: msg.content,
|
|
type: msg.type as "user" | "bot" | "system",
|
|
timestamp: msg.createdAt,
|
|
metadata: msg.metadata,
|
|
})),
|
|
userPreferences: session.context,
|
|
};
|
|
|
|
// Generate response using LangChain
|
|
const langChainResponse = await this.langChainService.generateResponse(userMessage.content, context);
|
|
|
|
// Create bot message
|
|
const botMessage = new ChatMessage();
|
|
botMessage.content = langChainResponse.message;
|
|
botMessage.type = MessageType.BOT;
|
|
botMessage.session = session;
|
|
botMessage.responseToId = userMessage.id;
|
|
botMessage.tokensUsed = langChainResponse.tokensUsed;
|
|
botMessage.metadata = {
|
|
confidence: langChainResponse.confidence,
|
|
sources: langChainResponse.sources,
|
|
llmContext: langChainResponse.context,
|
|
provider: "langchain",
|
|
documentsRetrieved: langChainResponse.context?.documentsRetrieved || 0,
|
|
};
|
|
|
|
await em.persistAndFlush(botMessage);
|
|
|
|
// Update session last message time
|
|
await this.chatSessionRepo.updateLastMessageTime(session.id, em);
|
|
|
|
this.logger.log(`Generated LangChain bot response for session ${session.id}`);
|
|
} catch (error) {
|
|
this.logger.error(`Failed to generate LangChain bot response for session ${session.id}`, error);
|
|
|
|
// Fallback to regular LLM service
|
|
this.logger.log(`Falling back to regular LLM service for session ${session.id}`);
|
|
await this.generateBotResponse(session, userMessage, userId);
|
|
}
|
|
}
|
|
|
|
//************************************ */
|
|
|
|
async refreshLangChainData() {
|
|
try {
|
|
await this.langChainService.refreshVectorStore();
|
|
this.logger.log("LangChain vector store refreshed successfully");
|
|
return { message: "LangChain vector store refreshed successfully" };
|
|
} catch (error) {
|
|
this.logger.error("Failed to refresh LangChain vector store", error);
|
|
throw new Error("Failed to refresh training data");
|
|
}
|
|
}
|
|
|
|
//************************************ */
|
|
}
|