package service import ( "context" "fmt" "strings" "github.com/gochat/gochat/internal/llm" "github.com/gochat/gochat/internal/model" "github.com/pgvector/pgvector-go" applogger "github.com/gochat/gochat/pkg/logger" ) // RAGService implements Retrieval-Augmented Generation for Captain knowledge base Q&A. // Reference: Chatwoot Captain::Llm::GraphService + Captain::Tools::CaptainKnowledgeBaseTool // // RAG flow: // 1. Take user question → generate embedding via LLM provider // 2. Search assistant response embeddings in pgvector for top-k relevant FAQs // 3. Build prompt with retrieved context + question → call LLM for answer // 4. Return answer with source document references // RAGQueryRequest holds the input for a RAG Q&A query. type RAGQueryRequest struct { AssistantID uint `json:"assistant_id" validate:"required"` Question string `json:"question" validate:"required,min=1"` TopK int `json:"top_k,omitempty"` // number of FAQ responses to retrieve, default 5 } // RAGQueryResult holds the answer from RAG Q&A. type RAGQueryResult struct { Answer string `json:"answer"` Sources []RAGSource `json:"sources,omitempty"` Confidence float64 `json:"confidence"` AssistantID uint `json:"assistant_id"` Conversation uint `json:"conversation_id,omitempty"` } // RAGSource references a source FAQ response used in the answer. type RAGSource struct { ResponseID uint `json:"response_id"` Question string `json:"question"` Answer string `json:"answer"` DocumentID uint `json:"document_id,omitempty"` Score float64 `json:"score"` } // AssistantRepoIface defines the repository interface RAGService needs from CaptainAssistantRepo. type AssistantRepoIface interface { GetByID(ctx context.Context, id uint) (*model.CaptainAssistant, error) } // ResponseRepoIface defines the repository interface RAGService needs from CaptainAssistantResponseRepo. type ResponseRepoIface interface { SearchByEmbedding(ctx context.Context, assistantID uint, embedding pgvector.Vector, limit int) ([]model.CaptainAssistantResponse, error) GetByID(ctx context.Context, id uint) (*model.CaptainAssistantResponse, error) Update(ctx context.Context, resp *model.CaptainAssistantResponse) error } // RAGService orchestrates embedding search + LLM generation for knowledge base Q&A. type RAGService struct { responseRepo ResponseRepoIface assistantRepo AssistantRepoIface llmProvider llm.Provider promptBuilder *SystemPromptBuilder } // NewRAGService creates a new RAGService. func NewRAGService( responseRepo ResponseRepoIface, assistantRepo AssistantRepoIface, llmProvider llm.Provider, ) *RAGService { return &RAGService{ responseRepo: responseRepo, assistantRepo: assistantRepo, llmProvider: llmProvider, promptBuilder: NewSystemPromptBuilder(), } } // Query performs a RAG Q&A: embed question → search FAQs → generate answer. func (s *RAGService) Query(ctx context.Context, accountID uint, req *RAGQueryRequest) (*RAGQueryResult, error) { if req.TopK <= 0 { req.TopK = 5 } // Step 1: Verify assistant exists and is active assistant, err := s.assistantRepo.GetByID(ctx, req.AssistantID) if err != nil { return nil, fmt.Errorf("assistant not found: %w", err) } if assistant.Status != model.AssistantStatusActive { return nil, fmt.Errorf("assistant is not active (status: %s)", assistant.Status) } // Get assistant config for prompt building and model selection cfg, err := assistant.GetConfig() if err != nil { applogger.L().Warnf("RAG get assistant config: %v", err) cfg = &model.AssistantConfig{} } // Step 2: Generate embedding for the question embedResp, err := s.llmProvider.CreateEmbedding(ctx, llm.EmbeddingRequest{ Model: "text-embedding-3-small", Input: []string{req.Question}, }) if err != nil { applogger.L().Errorf("RAG embedding generation: %v", err) return nil, fmt.Errorf("embedding generation failed: %w", err) } if len(embedResp.Data) == 0 { return nil, fmt.Errorf("no embedding returned for question") } questionEmbedding := embedResp.Data[0].Embedding // Convert float64 embedding to float32 for pgvector float32Emb := make([]float32, len(questionEmbedding)) for i, v := range questionEmbedding { float32Emb[i] = float32(v) } pgvectorEmb := pgvector.NewVector(float32Emb) // Step 3: Search assistant responses by embedding similarity (pgvector cosine distance) results, err := s.responseRepo.SearchByEmbedding(ctx, req.AssistantID, pgvectorEmb, req.TopK) if err != nil { applogger.L().Errorf("RAG response search: %v", err) return nil, fmt.Errorf("FAQ search failed: %w", err) } if len(results) == 0 { // No relevant FAQs found — fallback to LLM without context return s.queryWithoutContext(ctx, assistant, cfg, req.Question) } // Step 4: Build RAG context from retrieved FAQ responses var contextParts []string var sources []RAGSource for i, r := range results { contextParts = append(contextParts, fmt.Sprintf("[FAQ %d]\nQ: %s\nA: %s", i+1, r.Question, r.Answer)) // Estimate score (pgvector cosine distance is already sorted by closest) score := 1.0 - float64(i)*0.1 // rough heuristic since we don't have raw distance if i == 0 { score = 0.9 // top result is likely high-confidence } sources = append(sources, RAGSource{ ResponseID: r.ID, Question: r.Question, Answer: r.Answer, Score: score, }) if r.DocumentableID != nil { sources[len(sources)-1].DocumentID = *r.DocumentableID } } ragContext := strings.Join(contextParts, "\n\n") // Step 5: Build prompt with RAG context + system prompt systemPrompt := s.promptBuilder.BuildAssistantPrompt(assistant, cfg) systemPrompt += fmt.Sprintf("\n\nKnowledge Base Context:\n%s", ragContext) systemPrompt += "\n\nImportant: When answering, reference the FAQ entries by number. If the FAQ context doesn't contain enough information to answer, say so honestly." messages := []llm.ChatMessage{ {Role: "system", Content: systemPrompt}, {Role: "user", Content: req.Question}, } // Step 6: Call LLM for answer generation modelName := cfg.Model if modelName == "" { modelName = "gpt-4" } temperature := cfg.Temperature if temperature == 0 { temperature = 0.3 // lower temp for factual answers } llmResp, err := s.llmProvider.ChatCompletion(ctx, llm.ChatRequest{ Model: modelName, Messages: messages, Temperature: temperature, MaxTokens: 1024, }) if err != nil { applogger.L().Errorf("RAG LLM generation: %v", err) return nil, fmt.Errorf("answer generation failed: %w", err) } answer := "" confidence := 0.0 if len(llmResp.Choices) > 0 { answer = llmResp.Choices[0].Message.Content if len(sources) > 0 { confidence = sources[0].Score } } return &RAGQueryResult{ Answer: answer, Sources: sources, Confidence: confidence, AssistantID: req.AssistantID, }, nil } // queryWithoutContext handles the case when no relevant FAQ responses are found. // Falls back to generating an answer purely from the assistant's configuration. func (s *RAGService) queryWithoutContext(ctx context.Context, assistant *model.CaptainAssistant, cfg *model.AssistantConfig, question string) (*RAGQueryResult, error) { systemPrompt := s.promptBuilder.BuildAssistantPrompt(assistant, cfg) systemPrompt += "\n\nNote: No relevant FAQ entries were found for this question. Answer based on your general knowledge, but indicate that the answer may not be specific to the product." modelName := cfg.Model if modelName == "" { modelName = "gpt-4" } temperature := cfg.Temperature if temperature == 0 { temperature = 0.5 } llmResp, err := s.llmProvider.ChatCompletion(ctx, llm.ChatRequest{ Model: modelName, Messages: []llm.ChatMessage{ {Role: "system", Content: systemPrompt}, {Role: "user", Content: question}, }, Temperature: temperature, MaxTokens: 512, }) if err != nil { return nil, fmt.Errorf("fallback LLM generation failed: %w", err) } answer := "" if len(llmResp.Choices) > 0 { answer = llmResp.Choices[0].Message.Content } return &RAGQueryResult{ Answer: answer, Confidence: 0.3, // low confidence when no context AssistantID: assistant.ID, }, nil } // IndexResponse generates embeddings for an assistant response and stores them for search. // Called when a response is created or its question/answer is updated. func (s *RAGService) IndexResponse(ctx context.Context, responseID uint) error { resp, err := s.responseRepo.GetByID(ctx, responseID) if err != nil { return fmt.Errorf("response not found: %w", err) } // Generate embedding from question + answer for better semantic matching inputText := fmt.Sprintf("Q: %s\nA: %s", resp.Question, resp.Answer) embedResp, err := s.llmProvider.CreateEmbedding(ctx, llm.EmbeddingRequest{ Model: "text-embedding-3-small", Input: []string{inputText}, }) if err != nil { applogger.L().Errorf("RAG response embedding: %v", err) return fmt.Errorf("response embedding failed: %w", err) } if len(embedResp.Data) > 0 { float32Emb := make([]float32, len(embedResp.Data[0].Embedding)) for i, v := range embedResp.Data[0].Embedding { float32Emb[i] = float32(v) } resp.Embedding = pgvector.NewVector(float32Emb) if err := s.responseRepo.Update(ctx, resp); err != nil { applogger.L().Errorf("RAG update response embedding: %v", err) return fmt.Errorf("store embedding failed: %w", err) } } applogger.L().Infof("RAG indexed response %d for assistant %d", resp.ID, resp.AssistantID) return nil }