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 { 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 { 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"); } } //************************************ */ }