127 lines
4.0 KiB
Go
127 lines
4.0 KiB
Go
package kb
|
|
|
|
import (
|
|
"context"
|
|
"fmt"
|
|
"strings"
|
|
|
|
"ai-operator/internal/config"
|
|
"ai-operator/internal/embedding"
|
|
)
|
|
|
|
type Service struct {
|
|
repo Repository
|
|
embed embedding.Provider
|
|
cfg config.KBConfig
|
|
}
|
|
|
|
func NewService(repo Repository, embed embedding.Provider, cfg config.KBConfig) *Service {
|
|
return &Service{repo: repo, embed: embed, cfg: cfg}
|
|
}
|
|
|
|
type SearchResponse struct {
|
|
OK bool
|
|
ReasonCode, MessageKey string
|
|
Results []SearchResult
|
|
Citations []Citation
|
|
CrossLanguageFallbackUsed bool
|
|
}
|
|
|
|
func (s *Service) Search(ctx context.Context, req SearchRequest) (SearchResponse, error) {
|
|
q := strings.TrimSpace(req.Query)
|
|
if q == "" {
|
|
return SearchResponse{OK: false, ReasonCode: "query_too_short", MessageKey: "knowledge.query_too_short"}, nil
|
|
}
|
|
if s.cfg.QueryMaxChars > 0 && len([]rune(q)) > s.cfg.QueryMaxChars {
|
|
return SearchResponse{OK: false, ReasonCode: "query_too_long", MessageKey: "knowledge.query_too_long"}, nil
|
|
}
|
|
if req.Language != "ru" && req.Language != "kk" {
|
|
return SearchResponse{OK: false, ReasonCode: "language_required", MessageKey: "knowledge.search_denied_language"}, nil
|
|
}
|
|
if req.RegionCode == "" {
|
|
return SearchResponse{OK: false, ReasonCode: "region_required", MessageKey: "knowledge.search_denied_region"}, nil
|
|
}
|
|
if req.Limit <= 0 {
|
|
req.Limit = s.cfg.DefaultLimit
|
|
}
|
|
if req.Limit <= 0 {
|
|
req.Limit = 5
|
|
}
|
|
if s.cfg.MaxLimit > 0 && req.Limit > s.cfg.MaxLimit {
|
|
req.Limit = s.cfg.MaxLimit
|
|
}
|
|
if req.MinScore == 0 {
|
|
req.MinScore = s.cfg.MinScore
|
|
}
|
|
emb, err := s.embed.Embed(ctx, []string{q})
|
|
if err != nil {
|
|
return SearchResponse{OK: false, ReasonCode: "knowledge_base_unavailable", MessageKey: "knowledge.unavailable"}, err
|
|
}
|
|
req.Query = q
|
|
req.IncludeGlobal = true
|
|
results, err := s.repo.Search(ctx, req, emb[0])
|
|
if err != nil {
|
|
return SearchResponse{OK: false, ReasonCode: "knowledge_base_unavailable", MessageKey: "knowledge.unavailable"}, err
|
|
}
|
|
fallback := false
|
|
if len(results) == 0 && req.Language == "kk" && (req.CrossLanguageFallback || s.cfg.CrossLanguageFallback) {
|
|
ruReq := req
|
|
ruReq.Language = "ru"
|
|
results, err = s.repo.Search(ctx, ruReq, emb[0])
|
|
if err != nil {
|
|
return SearchResponse{OK: false, ReasonCode: "knowledge_base_unavailable", MessageKey: "knowledge.unavailable"}, err
|
|
}
|
|
if len(results) > 0 {
|
|
fallback = true
|
|
for i := range results {
|
|
results[i].CrossLanguageFallback = true
|
|
results[i].SourceLanguage = "ru"
|
|
}
|
|
}
|
|
}
|
|
if len(results) == 0 {
|
|
return SearchResponse{OK: false, ReasonCode: "no_relevant_knowledge", MessageKey: "knowledge.no_answer"}, nil
|
|
}
|
|
cites := make([]Citation, 0, len(results))
|
|
for _, r := range results {
|
|
cites = append(cites, r.Citation)
|
|
}
|
|
return SearchResponse{OK: true, ReasonCode: "ok", MessageKey: "knowledge.results_found", Results: results, Citations: cites, CrossLanguageFallbackUsed: fallback}, nil
|
|
}
|
|
|
|
func (s *Service) Health(ctx context.Context) (Health, error) { return s.repo.Health(ctx) }
|
|
|
|
func (s *Service) IngestJSONL(ctx context.Context, path string, allowDisabled bool) (IngestResult, error) {
|
|
docs, parseErrs, err := LoadJSONLDocuments(ctx, path, allowDisabled)
|
|
if err != nil {
|
|
return IngestResult{}, err
|
|
}
|
|
res := IngestResult{DocsSeen: len(docs) + len(parseErrs), Errors: parseErrs}
|
|
for _, doc := range docs {
|
|
chunks := ChunkDocument(doc, DefaultChunkerConfig())
|
|
texts := make([]string, len(chunks))
|
|
for i, ch := range chunks {
|
|
texts[i] = ch.Content
|
|
}
|
|
vecs, err := s.embed.Embed(ctx, texts)
|
|
if err != nil {
|
|
res.Errors = append(res.Errors, fmt.Sprintf("%s: embedding failed", doc.ExternalID))
|
|
res.Skipped++
|
|
continue
|
|
}
|
|
for i := range chunks {
|
|
chunks[i].Embedding = vecs[i]
|
|
chunks[i].EmbeddingModel = s.embed.Model()
|
|
chunks[i].EmbeddingProvider = s.embed.ProviderName()
|
|
}
|
|
if err := s.repo.UpsertDocument(ctx, doc, chunks); err != nil {
|
|
res.Errors = append(res.Errors, fmt.Sprintf("%s: %v", doc.ExternalID, err))
|
|
res.Skipped++
|
|
continue
|
|
}
|
|
res.DocsIngested++
|
|
res.ChunksCreated += len(chunks)
|
|
}
|
|
return res, nil
|
|
}
|