mirror of
https://github.com/Tencent/WeKnora.git
synced 2026-08-30 16:53:21 +08:00
Revert "refactor(knowledge): split 9.8k-line service file by responsibility"
This reverts commit f67d0b53ce.
This commit is contained in:
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -1,686 +0,0 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Tencent/WeKnora/internal/application/service/retriever"
|
||||
"github.com/Tencent/WeKnora/internal/logger"
|
||||
"github.com/Tencent/WeKnora/internal/types"
|
||||
"github.com/Tencent/WeKnora/internal/types/interfaces"
|
||||
"github.com/hibiken/asynq"
|
||||
"golang.org/x/sync/errgroup"
|
||||
)
|
||||
|
||||
// collectImageURLs extracts unique provider:// image URLs from image_info JSON strings.
|
||||
func collectImageURLs(ctx context.Context, imageInfos []string) []string {
|
||||
seen := make(map[string]struct{})
|
||||
var urls []string
|
||||
for _, info := range imageInfos {
|
||||
if info == "" {
|
||||
continue
|
||||
}
|
||||
var images []*types.ImageInfo
|
||||
if err := json.Unmarshal([]byte(info), &images); err != nil {
|
||||
logger.Warnf(ctx, "Failed to parse image_info JSON: %v", err)
|
||||
continue
|
||||
}
|
||||
for _, img := range images {
|
||||
if img.URL != "" {
|
||||
if _, exists := seen[img.URL]; !exists {
|
||||
seen[img.URL] = struct{}{}
|
||||
urls = append(urls, img.URL)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
return urls
|
||||
}
|
||||
|
||||
// deleteExtractedImages deletes all extracted image files from storage.
|
||||
// Standalone function — callable from both knowledgeService and knowledgeBaseService.
|
||||
// Errors are logged but do not fail the overall deletion.
|
||||
func deleteExtractedImages(ctx context.Context, fileSvc interfaces.FileService, imageURLs []string) {
|
||||
if len(imageURLs) == 0 {
|
||||
return
|
||||
}
|
||||
logger.Infof(ctx, "Deleting %d extracted images", len(imageURLs))
|
||||
for _, url := range imageURLs {
|
||||
if err := fileSvc.DeleteFile(ctx, url); err != nil {
|
||||
logger.Errorf(ctx, "Failed to delete extracted image %s: %v", url, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// DeleteKnowledge deletes a knowledge entry and all related resources
|
||||
func (s *knowledgeService) DeleteKnowledge(ctx context.Context, id string) error {
|
||||
// Get the knowledge entry
|
||||
knowledge, err := s.repo.GetKnowledgeByID(ctx, ctx.Value(types.TenantIDContextKey).(uint64), id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Mark as deleting first to prevent async task conflicts
|
||||
// This ensures that any running async tasks will detect the deletion and abort
|
||||
originalStatus := knowledge.ParseStatus
|
||||
knowledge.ParseStatus = types.ParseStatusDeleting
|
||||
knowledge.UpdatedAt = time.Now()
|
||||
if err := s.repo.UpdateKnowledge(ctx, knowledge); err != nil {
|
||||
logger.GetLogger(ctx).WithField("error", err).Errorf("DeleteKnowledge failed to mark as deleting")
|
||||
// Continue with deletion even if marking fails
|
||||
} else {
|
||||
logger.Infof(ctx, "Marked knowledge %s as deleting (previous status: %s)", id, originalStatus)
|
||||
}
|
||||
|
||||
// Resolve file service for this KB before spawning goroutines
|
||||
kb, _ := s.kbService.GetKnowledgeBaseByID(ctx, knowledge.KnowledgeBaseID)
|
||||
kbFileSvc := s.resolveFileService(ctx, kb)
|
||||
|
||||
// Collect image URLs before chunks are deleted (ImageInfo references are lost after deletion)
|
||||
tenantID := ctx.Value(types.TenantIDContextKey).(uint64)
|
||||
chunkImageInfos, err := s.chunkService.GetRepository().ListImageInfoByKnowledgeIDs(ctx, tenantID, []string{id})
|
||||
if err != nil {
|
||||
logger.Errorf(ctx, "Failed to collect image URLs for cleanup: %v", err)
|
||||
}
|
||||
var imageInfoStrs []string
|
||||
for _, ci := range chunkImageInfos {
|
||||
imageInfoStrs = append(imageInfoStrs, ci.ImageInfo)
|
||||
}
|
||||
imageURLs := collectImageURLs(ctx, imageInfoStrs)
|
||||
|
||||
wg := errgroup.Group{}
|
||||
// Delete knowledge embeddings from vector store.
|
||||
// Skip entirely when the knowledge has no embedding model (e.g. Wiki-only KB):
|
||||
// nothing was ever written to the vector store, so there is nothing to delete,
|
||||
// and GetEmbeddingModel would fail with "model ID cannot be empty".
|
||||
if strings.TrimSpace(knowledge.EmbeddingModelID) != "" {
|
||||
wg.Go(func() error {
|
||||
tenantInfo := ctx.Value(types.TenantInfoContextKey).(*types.Tenant)
|
||||
retrieveEngine, err := retriever.NewCompositeRetrieveEngine(
|
||||
s.retrieveEngine,
|
||||
tenantInfo.GetEffectiveEngines(),
|
||||
)
|
||||
if err != nil {
|
||||
logger.GetLogger(ctx).WithField("error", err).Errorf("DeleteKnowledge delete knowledge embedding failed")
|
||||
return err
|
||||
}
|
||||
embeddingModel, err := s.modelService.GetEmbeddingModel(ctx, knowledge.EmbeddingModelID)
|
||||
if err != nil {
|
||||
logger.GetLogger(ctx).WithField("error", err).Errorf("DeleteKnowledge delete knowledge embedding failed")
|
||||
return err
|
||||
}
|
||||
if err := retrieveEngine.DeleteByKnowledgeIDList(ctx, []string{knowledge.ID}, embeddingModel.GetDimensions(), knowledge.Type); err != nil {
|
||||
logger.GetLogger(ctx).WithField("error", err).Errorf("DeleteKnowledge delete knowledge embedding failed")
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
})
|
||||
} else {
|
||||
logger.Infof(ctx, "Knowledge %s has no embedding model, skipping vector store cleanup", knowledge.ID)
|
||||
}
|
||||
|
||||
// Delete all chunks associated with this knowledge
|
||||
wg.Go(func() error {
|
||||
if err := s.chunkService.DeleteChunksByKnowledgeID(ctx, knowledge.ID); err != nil {
|
||||
logger.GetLogger(ctx).WithField("error", err).Errorf("DeleteKnowledge delete chunks failed")
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
})
|
||||
|
||||
// Delete the physical file and extracted images if they exist
|
||||
wg.Go(func() error {
|
||||
if knowledge.FilePath != "" {
|
||||
if err := kbFileSvc.DeleteFile(ctx, knowledge.FilePath); err != nil {
|
||||
logger.GetLogger(ctx).WithField("error", err).Errorf("DeleteKnowledge delete file failed")
|
||||
}
|
||||
}
|
||||
deleteExtractedImages(ctx, kbFileSvc, imageURLs)
|
||||
tenantInfo := ctx.Value(types.TenantInfoContextKey).(*types.Tenant)
|
||||
tenantInfo.StorageUsed -= knowledge.StorageSize
|
||||
if err := s.tenantRepo.AdjustStorageUsed(ctx, tenantInfo.ID, -knowledge.StorageSize); err != nil {
|
||||
logger.GetLogger(ctx).WithField("error", err).Errorf("DeleteKnowledge update tenant storage used failed")
|
||||
}
|
||||
return nil
|
||||
})
|
||||
|
||||
// Delete the knowledge graph
|
||||
wg.Go(func() error {
|
||||
namespace := types.NameSpace{KnowledgeBase: knowledge.KnowledgeBaseID, Knowledge: knowledge.ID}
|
||||
if err := s.graphEngine.DelGraph(ctx, []types.NameSpace{namespace}); err != nil {
|
||||
logger.GetLogger(ctx).WithField("error", err).Errorf("DeleteKnowledge delete knowledge graph failed")
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
})
|
||||
|
||||
// Clean up wiki pages that reference this knowledge. Pass the full
|
||||
// knowledge object so cleanup can source title/summary from the row
|
||||
// itself rather than reaching into possibly-not-yet-written wiki pages.
|
||||
if kb != nil && kb.IsWikiEnabled() {
|
||||
wg.Go(func() error {
|
||||
s.cleanupWikiOnKnowledgeDelete(ctx, knowledge)
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
if err = wg.Wait(); err != nil {
|
||||
return err
|
||||
}
|
||||
// Delete the knowledge entry itself from the database
|
||||
return s.repo.DeleteKnowledge(ctx, ctx.Value(types.TenantIDContextKey).(uint64), id)
|
||||
}
|
||||
|
||||
// cleanupWikiOnKnowledgeDelete handles wiki pages when a source document is deleted.
|
||||
//
|
||||
// There are three sources of truth we must keep consistent:
|
||||
// - The knowledge row (being soft-deleted right now by the caller)
|
||||
// - Wiki pages whose source_refs include this knowledge
|
||||
// - Pending/in-flight wiki_ingest tasks that may create *new* pages pointing at it
|
||||
//
|
||||
// The function is deliberately best-effort and idempotent:
|
||||
// - It writes a tombstone + scrubs pending ingest ops so new pages cannot be
|
||||
// born with a stale source_ref (guards (a) queued ingest and (b) ingest
|
||||
// tasks mid-LLM call — both consult the tombstone before writing).
|
||||
// - It immediately reconciles any pages already present (delete-if-only-ref
|
||||
// or strip-ref-if-multi).
|
||||
// - It *unconditionally* enqueues a retract task. Crucially we DO NOT gate
|
||||
// enqueue on "pages currently exist": in the ingest/delete race the
|
||||
// knowledge may have pages that exist only after this function returns
|
||||
// (the ingest task fires later and, absent the tombstone, would have
|
||||
// created them). The retract handler re-queries ListPagesBySourceRef at
|
||||
// run time, so even with an empty PageSlugs it will do the right thing —
|
||||
// and at worst it's a cheap no-op.
|
||||
func (s *knowledgeService) cleanupWikiOnKnowledgeDelete(ctx context.Context, knowledge *types.Knowledge) {
|
||||
if knowledge == nil {
|
||||
return
|
||||
}
|
||||
kbID := knowledge.KnowledgeBaseID
|
||||
knowledgeID := knowledge.ID
|
||||
if kbID == "" || knowledgeID == "" {
|
||||
return
|
||||
}
|
||||
|
||||
// (1) Tombstone + scrub pending ingest — must happen first so any
|
||||
// wiki_ingest task that wakes up between here and the retract enqueue
|
||||
// below sees "knowledge gone" and bails out.
|
||||
s.markKnowledgeDeletedForWiki(ctx, kbID, knowledgeID)
|
||||
s.scrubWikiPendingIngest(ctx, kbID, knowledgeID, "cleanup")
|
||||
|
||||
// Pull title/summary from the knowledge itself — do NOT read them from
|
||||
// existing wiki pages. In the race window wiki pages may not exist yet,
|
||||
// and even when they do their "summary" is the LLM-extracted one which
|
||||
// we're about to invalidate anyway. The knowledge row still has the
|
||||
// original Title/FileName/Description, which is what the retract prompt
|
||||
// actually wants.
|
||||
docTitle := knowledge.Title
|
||||
if docTitle == "" {
|
||||
docTitle = knowledge.FileName
|
||||
}
|
||||
if docTitle == "" {
|
||||
docTitle = knowledgeID
|
||||
}
|
||||
docSummary := knowledge.Description
|
||||
|
||||
// (2) Immediate reconciliation for pages already present. If ingest
|
||||
// hasn't run yet this simply finds nothing; that's fine — see (3).
|
||||
pages, err := s.wikiRepo.ListBySourceRef(ctx, kbID, knowledgeID)
|
||||
if err != nil {
|
||||
logger.Warnf(ctx, "wiki cleanup: failed to list pages by source ref %s: %v", knowledgeID, err)
|
||||
pages = nil
|
||||
}
|
||||
|
||||
// Prefer the on-disk summary if the summary page already exists (it's
|
||||
// richer than the raw user-provided description). Leave docSummary
|
||||
// untouched otherwise so we still pass something meaningful downstream.
|
||||
for _, page := range pages {
|
||||
if page.PageType == types.WikiPageTypeSummary && page.Summary != "" {
|
||||
docSummary = page.Summary
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
var deletedSlugs []string
|
||||
var retractSlugs []string
|
||||
for _, page := range pages {
|
||||
if page.PageType == types.WikiPageTypeIndex || page.PageType == types.WikiPageTypeLog {
|
||||
continue
|
||||
}
|
||||
|
||||
remaining := removeSourceRef(page.SourceRefs, knowledgeID)
|
||||
|
||||
if len(remaining) == 0 {
|
||||
if err := s.wikiService.DeletePage(ctx, kbID, page.Slug); err != nil {
|
||||
logger.Warnf(ctx, "wiki cleanup: failed to delete page %s: %v", page.Slug, err)
|
||||
} else {
|
||||
deletedSlugs = append(deletedSlugs, page.Slug)
|
||||
}
|
||||
} else {
|
||||
page.SourceRefs = remaining
|
||||
if err := s.wikiService.UpdatePageMeta(ctx, page); err != nil {
|
||||
logger.Warnf(ctx, "wiki cleanup: failed to update source refs for page %s: %v", page.Slug, err)
|
||||
} else {
|
||||
retractSlugs = append(retractSlugs, page.Slug)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if len(deletedSlugs) > 0 {
|
||||
logger.Infof(ctx, "wiki cleanup: deleted %d pages after knowledge %s deletion: %v",
|
||||
len(deletedSlugs), knowledgeID, deletedSlugs)
|
||||
}
|
||||
|
||||
allAffectedSlugs := append(retractSlugs, deletedSlugs...)
|
||||
|
||||
// (3) Unconditionally enqueue the retract task. See function comment —
|
||||
// an empty PageSlugs is not a bug, it's the signal "re-query at run
|
||||
// time". The handler will ListPagesBySourceRef again, pick up any
|
||||
// pages that materialised after we looked, and also rebuild index/log
|
||||
// so the knowledge's disappearance is reflected in the UI.
|
||||
lang, _ := types.LanguageFromContext(ctx)
|
||||
tenantID, _ := types.TenantIDFromContext(ctx)
|
||||
EnqueueWikiRetract(ctx, s.task, s.redisClient, WikiRetractPayload{
|
||||
TenantID: tenantID,
|
||||
KnowledgeBaseID: kbID,
|
||||
KnowledgeID: knowledgeID,
|
||||
DocTitle: docTitle,
|
||||
DocSummary: docSummary,
|
||||
Language: lang,
|
||||
PageSlugs: allAffectedSlugs,
|
||||
})
|
||||
logger.Infof(ctx, "wiki cleanup: enqueued retract task for knowledge %s (%d known slugs: %v)",
|
||||
knowledgeID, len(allAffectedSlugs), allAffectedSlugs)
|
||||
}
|
||||
|
||||
// markKnowledgeDeletedForWiki writes a short-TTL tombstone so any wiki_ingest
|
||||
// task still running or queued for this knowledge can short-circuit before
|
||||
// resurrecting a page with a stale source_ref. No-op when Redis is absent.
|
||||
func (s *knowledgeService) markKnowledgeDeletedForWiki(ctx context.Context, kbID, knowledgeID string) {
|
||||
if s.redisClient == nil || kbID == "" || knowledgeID == "" {
|
||||
return
|
||||
}
|
||||
key := WikiDeletedTombstoneKey(kbID, knowledgeID)
|
||||
if err := s.redisClient.Set(ctx, key, "1", wikiDeletedTTL).Err(); err != nil {
|
||||
logger.Warnf(ctx, "wiki cleanup: failed to write tombstone %s: %v", key, err)
|
||||
}
|
||||
}
|
||||
|
||||
// scrubWikiPendingIngest removes queued WikiOpIngest entries for a knowledge
|
||||
// from the debounced pending list. Used by both the delete path (we're about
|
||||
// to soft-delete the doc, no point ingesting it) and the reparse path (the
|
||||
// old chunks are about to vanish, so any pending ingest would either race
|
||||
// with the cleanup or no-op on an empty chunk set — and the post-process
|
||||
// task will enqueue a fresh ingest once new chunks land anyway).
|
||||
//
|
||||
// Retract entries stay put — delete still needs them to unlink referencing
|
||||
// pages, and reparse never enqueues retracts for the doc being reparsed.
|
||||
//
|
||||
// We use LREM against JSON-encoded entries plus a best-effort raw-UUID
|
||||
// fallback for backward compatibility with the legacy format documented in
|
||||
// peekPendingList.
|
||||
func (s *knowledgeService) scrubWikiPendingIngest(ctx context.Context, kbID, knowledgeID, reason string) {
|
||||
if s.redisClient == nil || kbID == "" || knowledgeID == "" {
|
||||
return
|
||||
}
|
||||
pendingKey := wikiPendingKeyPrefix + kbID
|
||||
|
||||
// Best-effort: inspect the list, remove matching ingest entries one by one.
|
||||
// The list is bounded (wikiMaxDocsPerBatch at a time on the consumer
|
||||
// side, practical uploads rarely exceed a few dozen), so a single LRange
|
||||
// is safe.
|
||||
items, err := s.redisClient.LRange(ctx, pendingKey, 0, -1).Result()
|
||||
if err != nil {
|
||||
logger.Warnf(ctx, "wiki %s: failed to read pending list %s: %v", reason, pendingKey, err)
|
||||
return
|
||||
}
|
||||
removed := 0
|
||||
for _, item := range items {
|
||||
// Legacy raw-UUID form
|
||||
if item == knowledgeID {
|
||||
if n, err := s.redisClient.LRem(ctx, pendingKey, 0, item).Result(); err == nil {
|
||||
removed += int(n)
|
||||
}
|
||||
continue
|
||||
}
|
||||
if !strings.HasPrefix(item, "{") {
|
||||
continue
|
||||
}
|
||||
var op WikiPendingOp
|
||||
if err := json.Unmarshal([]byte(item), &op); err != nil {
|
||||
continue
|
||||
}
|
||||
if op.KnowledgeID != knowledgeID || op.Op != WikiOpIngest {
|
||||
continue
|
||||
}
|
||||
if n, err := s.redisClient.LRem(ctx, pendingKey, 0, item).Result(); err == nil {
|
||||
removed += int(n)
|
||||
}
|
||||
}
|
||||
if removed > 0 {
|
||||
logger.Infof(ctx, "wiki %s: scrubbed %d pending ingest ops for knowledge %s", reason, removed, knowledgeID)
|
||||
}
|
||||
}
|
||||
|
||||
// prepareWikiForReparse is the reparse counterpart to
|
||||
// cleanupWikiOnKnowledgeDelete. It aligns reparse with the same "pending
|
||||
// queue hygiene" the delete path already enforces, without taking any
|
||||
// destructive action against existing pages.
|
||||
//
|
||||
// Why no retract / tombstone here: reparse is not a "K is gone" event, it's
|
||||
// a "K's contribution is about to be swapped for a new version" event. The
|
||||
// actual swap happens asynchronously inside mapOneDocument (see its
|
||||
// oldPageSlugs handling) — that's where we have both the old page set and
|
||||
// the freshly extracted candidate slugs, which is exactly the information
|
||||
// the WikiPageModifyPrompt needs to do a correct replace-not-append.
|
||||
//
|
||||
// So the only thing worth doing synchronously at reparse time is keeping
|
||||
// the Redis pending list clean so the re-ingest enqueued by
|
||||
// KnowledgePostProcess doesn't race with a stale ingest op that would
|
||||
// fire mid-flight against zero chunks.
|
||||
func (s *knowledgeService) prepareWikiForReparse(ctx context.Context, knowledge *types.Knowledge) {
|
||||
if knowledge == nil {
|
||||
return
|
||||
}
|
||||
kbID := knowledge.KnowledgeBaseID
|
||||
knowledgeID := knowledge.ID
|
||||
if kbID == "" || knowledgeID == "" {
|
||||
return
|
||||
}
|
||||
s.scrubWikiPendingIngest(ctx, kbID, knowledgeID, "reparse")
|
||||
}
|
||||
|
||||
// removeSourceRef removes entries from source_refs that match a knowledge ID.
|
||||
// Handles both old format ("knowledgeID") and new format ("knowledgeID|title").
|
||||
func removeSourceRef(refs types.StringArray, knowledgeID string) types.StringArray {
|
||||
var result types.StringArray
|
||||
prefix := knowledgeID + "|"
|
||||
for _, ref := range refs {
|
||||
if ref == knowledgeID || strings.HasPrefix(ref, prefix) {
|
||||
continue
|
||||
}
|
||||
result = append(result, ref)
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// DeleteKnowledgeList deletes a knowledge entry and all related resources
|
||||
func (s *knowledgeService) DeleteKnowledgeList(ctx context.Context, ids []string) error {
|
||||
if len(ids) == 0 {
|
||||
return nil
|
||||
}
|
||||
// 1. Get the knowledge entry
|
||||
tenantInfo := ctx.Value(types.TenantInfoContextKey).(*types.Tenant)
|
||||
knowledgeList, err := s.repo.GetKnowledgeBatch(ctx, tenantInfo.ID, ids)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Mark all as deleting first to prevent async task conflicts
|
||||
for _, knowledge := range knowledgeList {
|
||||
knowledge.ParseStatus = types.ParseStatusDeleting
|
||||
knowledge.UpdatedAt = time.Now()
|
||||
if err := s.repo.UpdateKnowledge(ctx, knowledge); err != nil {
|
||||
logger.GetLogger(ctx).WithField("error", err).WithField("knowledge_id", knowledge.ID).
|
||||
Errorf("DeleteKnowledgeList failed to mark as deleting")
|
||||
// Continue with deletion even if marking fails
|
||||
}
|
||||
}
|
||||
logger.Infof(ctx, "Marked %d knowledge entries as deleting", len(knowledgeList))
|
||||
|
||||
// Pre-resolve file services per KB so goroutines don't need DB access
|
||||
kbFileServices := make(map[string]interfaces.FileService)
|
||||
for _, knowledge := range knowledgeList {
|
||||
if _, ok := kbFileServices[knowledge.KnowledgeBaseID]; !ok {
|
||||
kb, _ := s.kbService.GetKnowledgeBaseByID(ctx, knowledge.KnowledgeBaseID)
|
||||
kbFileServices[knowledge.KnowledgeBaseID] = s.resolveFileService(ctx, kb)
|
||||
}
|
||||
}
|
||||
|
||||
// Collect image URLs before chunks are deleted
|
||||
chunkImageInfos, err := s.chunkService.GetRepository().ListImageInfoByKnowledgeIDs(ctx, tenantInfo.ID, ids)
|
||||
if err != nil {
|
||||
logger.Errorf(ctx, "Failed to collect image URLs for batch cleanup: %v", err)
|
||||
}
|
||||
knowledgeToKB := make(map[string]string)
|
||||
for _, k := range knowledgeList {
|
||||
knowledgeToKB[k.ID] = k.KnowledgeBaseID
|
||||
}
|
||||
kbImageInfos := make(map[string][]string) // kbID → []imageInfo JSON
|
||||
for _, ci := range chunkImageInfos {
|
||||
kbID := knowledgeToKB[ci.KnowledgeID]
|
||||
kbImageInfos[kbID] = append(kbImageInfos[kbID], ci.ImageInfo)
|
||||
}
|
||||
kbImageURLs := make(map[string][]string) // kbID → []imageURL (deduplicated)
|
||||
for kbID, infos := range kbImageInfos {
|
||||
kbImageURLs[kbID] = collectImageURLs(ctx, infos)
|
||||
}
|
||||
|
||||
wg := errgroup.Group{}
|
||||
// 2. Delete knowledge embeddings from vector store
|
||||
wg.Go(func() error {
|
||||
tenantInfo := ctx.Value(types.TenantInfoContextKey).(*types.Tenant)
|
||||
retrieveEngine, err := retriever.NewCompositeRetrieveEngine(
|
||||
s.retrieveEngine,
|
||||
tenantInfo.GetEffectiveEngines(),
|
||||
)
|
||||
if err != nil {
|
||||
logger.GetLogger(ctx).WithField("error", err).Errorf("DeleteKnowledge delete knowledge embedding failed")
|
||||
return err
|
||||
}
|
||||
// Group by EmbeddingModelID and Type
|
||||
type groupKey struct {
|
||||
EmbeddingModelID string
|
||||
Type string
|
||||
}
|
||||
group := map[groupKey][]string{}
|
||||
for _, knowledge := range knowledgeList {
|
||||
key := groupKey{EmbeddingModelID: knowledge.EmbeddingModelID, Type: knowledge.Type}
|
||||
group[key] = append(group[key], knowledge.ID)
|
||||
}
|
||||
for key, knowledgeIDs := range group {
|
||||
// Wiki-only knowledge never had embeddings written to the vector store,
|
||||
// and its EmbeddingModelID is intentionally empty. Skip the whole group
|
||||
// to avoid the spurious "model ID cannot be empty" failure.
|
||||
if strings.TrimSpace(key.EmbeddingModelID) == "" {
|
||||
logger.Infof(ctx, "Skipping vector store cleanup for %d knowledge entries without embedding model", len(knowledgeIDs))
|
||||
continue
|
||||
}
|
||||
embeddingModel, err := s.modelService.GetEmbeddingModel(ctx, key.EmbeddingModelID)
|
||||
if err != nil {
|
||||
logger.GetLogger(ctx).WithField("error", err).Errorf("DeleteKnowledge get embedding model failed")
|
||||
return err
|
||||
}
|
||||
if err := retrieveEngine.DeleteByKnowledgeIDList(ctx, knowledgeIDs, embeddingModel.GetDimensions(), key.Type); err != nil {
|
||||
logger.GetLogger(ctx).
|
||||
WithField("error", err).
|
||||
Errorf("DeleteKnowledge delete knowledge embedding failed")
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
|
||||
// 3. Delete all chunks associated with this knowledge
|
||||
wg.Go(func() error {
|
||||
if err := s.chunkService.DeleteByKnowledgeList(ctx, ids); err != nil {
|
||||
logger.GetLogger(ctx).WithField("error", err).Errorf("DeleteKnowledge delete chunks failed")
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
})
|
||||
|
||||
// 4. Delete the physical file and extracted images if they exist
|
||||
wg.Go(func() error {
|
||||
storageAdjust := int64(0)
|
||||
for _, knowledge := range knowledgeList {
|
||||
if knowledge.FilePath != "" {
|
||||
fSvc := kbFileServices[knowledge.KnowledgeBaseID]
|
||||
if err := fSvc.DeleteFile(ctx, knowledge.FilePath); err != nil {
|
||||
logger.GetLogger(ctx).WithField("error", err).Errorf("DeleteKnowledge delete file failed")
|
||||
}
|
||||
}
|
||||
storageAdjust -= knowledge.StorageSize
|
||||
}
|
||||
// Delete extracted images per KB
|
||||
for kbID, urls := range kbImageURLs {
|
||||
fSvc := kbFileServices[kbID]
|
||||
if fSvc == nil {
|
||||
logger.Warnf(ctx, "No file service for KB %s, skipping %d image deletions", kbID, len(urls))
|
||||
continue
|
||||
}
|
||||
deleteExtractedImages(ctx, fSvc, urls)
|
||||
}
|
||||
tenantInfo.StorageUsed += storageAdjust
|
||||
if err := s.tenantRepo.AdjustStorageUsed(ctx, tenantInfo.ID, storageAdjust); err != nil {
|
||||
logger.GetLogger(ctx).WithField("error", err).Errorf("DeleteKnowledge update tenant storage used failed")
|
||||
}
|
||||
return nil
|
||||
})
|
||||
|
||||
// Delete the knowledge graph
|
||||
wg.Go(func() error {
|
||||
namespaces := []types.NameSpace{}
|
||||
for _, knowledge := range knowledgeList {
|
||||
namespaces = append(
|
||||
namespaces,
|
||||
types.NameSpace{KnowledgeBase: knowledge.KnowledgeBaseID, Knowledge: knowledge.ID},
|
||||
)
|
||||
}
|
||||
if err := s.graphEngine.DelGraph(ctx, namespaces); err != nil {
|
||||
logger.GetLogger(ctx).WithField("error", err).Errorf("DeleteKnowledge delete knowledge graph failed")
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
})
|
||||
|
||||
// Clean up wiki pages that reference deleted knowledge. cleanup needs
|
||||
// the full knowledge object (Title / Description) so the retract prompt
|
||||
// can describe the vanished document even when wiki pages haven't been
|
||||
// ingested yet — which is common in the batch-delete-shortly-after-upload
|
||||
// flow.
|
||||
wg.Go(func() error {
|
||||
for _, knowledge := range knowledgeList {
|
||||
kb, _ := s.kbService.GetKnowledgeBaseByID(ctx, knowledge.KnowledgeBaseID)
|
||||
if kb != nil && kb.IsWikiEnabled() {
|
||||
s.cleanupWikiOnKnowledgeDelete(ctx, knowledge)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
|
||||
if err = wg.Wait(); err != nil {
|
||||
return err
|
||||
}
|
||||
// 5. Delete the knowledge entry itself from the database
|
||||
return s.repo.DeleteKnowledgeList(ctx, tenantInfo.ID, ids)
|
||||
}
|
||||
|
||||
func (s *knowledgeService) cleanupKnowledgeResources(ctx context.Context, knowledge *types.Knowledge) error {
|
||||
logger.GetLogger(ctx).Infof("Cleaning knowledge resources before manual update, knowledge ID: %s", knowledge.ID)
|
||||
|
||||
var cleanupErr error
|
||||
|
||||
if knowledge.ParseStatus == types.ManualKnowledgeStatusDraft && knowledge.StorageSize == 0 {
|
||||
// Draft without indexed data, skip cleanup.
|
||||
return nil
|
||||
}
|
||||
|
||||
tenantInfo := ctx.Value(types.TenantInfoContextKey).(*types.Tenant)
|
||||
if knowledge.EmbeddingModelID != "" {
|
||||
retrieveEngine, err := retriever.NewCompositeRetrieveEngine(
|
||||
s.retrieveEngine,
|
||||
tenantInfo.GetEffectiveEngines(),
|
||||
)
|
||||
if err != nil {
|
||||
logger.GetLogger(ctx).WithField("error", err).Error("Failed to init retrieve engine during cleanup")
|
||||
cleanupErr = errors.Join(cleanupErr, err)
|
||||
} else {
|
||||
embeddingModel, modelErr := s.modelService.GetEmbeddingModel(ctx, knowledge.EmbeddingModelID)
|
||||
if modelErr != nil {
|
||||
logger.GetLogger(ctx).WithField("error", modelErr).Error("Failed to get embedding model during cleanup")
|
||||
cleanupErr = errors.Join(cleanupErr, modelErr)
|
||||
} else {
|
||||
if err := retrieveEngine.DeleteByKnowledgeIDList(ctx, []string{knowledge.ID}, embeddingModel.GetDimensions(), knowledge.Type); err != nil {
|
||||
logger.GetLogger(ctx).WithField("error", err).Error("Failed to delete manual knowledge index")
|
||||
cleanupErr = errors.Join(cleanupErr, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Collect image URLs before chunks are deleted
|
||||
kb, _ := s.kbService.GetKnowledgeBaseByID(ctx, knowledge.KnowledgeBaseID)
|
||||
fileSvc := s.resolveFileService(ctx, kb)
|
||||
chunkImageInfos, imgErr := s.chunkService.GetRepository().ListImageInfoByKnowledgeIDs(ctx, tenantInfo.ID, []string{knowledge.ID})
|
||||
if imgErr != nil {
|
||||
logger.GetLogger(ctx).WithField("error", imgErr).Error("Failed to collect image URLs for cleanup")
|
||||
cleanupErr = errors.Join(cleanupErr, imgErr)
|
||||
}
|
||||
var imageInfoStrs []string
|
||||
for _, ci := range chunkImageInfos {
|
||||
imageInfoStrs = append(imageInfoStrs, ci.ImageInfo)
|
||||
}
|
||||
imageURLs := collectImageURLs(ctx, imageInfoStrs)
|
||||
|
||||
if err := s.chunkService.DeleteChunksByKnowledgeID(ctx, knowledge.ID); err != nil {
|
||||
logger.GetLogger(ctx).WithField("error", err).Error("Failed to delete manual knowledge chunks")
|
||||
cleanupErr = errors.Join(cleanupErr, err)
|
||||
}
|
||||
|
||||
// Delete extracted images after chunks are deleted
|
||||
deleteExtractedImages(ctx, fileSvc, imageURLs)
|
||||
|
||||
namespace := types.NameSpace{KnowledgeBase: knowledge.KnowledgeBaseID, Knowledge: knowledge.ID}
|
||||
if err := s.graphEngine.DelGraph(ctx, []types.NameSpace{namespace}); err != nil {
|
||||
logger.GetLogger(ctx).WithField("error", err).Error("Failed to delete manual knowledge graph data")
|
||||
cleanupErr = errors.Join(cleanupErr, err)
|
||||
}
|
||||
|
||||
if knowledge.StorageSize > 0 {
|
||||
tenantInfo.StorageUsed -= knowledge.StorageSize
|
||||
if tenantInfo.StorageUsed < 0 {
|
||||
tenantInfo.StorageUsed = 0
|
||||
}
|
||||
if err := s.tenantRepo.AdjustStorageUsed(ctx, tenantInfo.ID, -knowledge.StorageSize); err != nil {
|
||||
logger.GetLogger(ctx).WithField("error", err).Error("Failed to adjust storage usage during manual cleanup")
|
||||
cleanupErr = errors.Join(cleanupErr, err)
|
||||
}
|
||||
knowledge.StorageSize = 0
|
||||
}
|
||||
|
||||
return cleanupErr
|
||||
}
|
||||
|
||||
// ProcessKnowledgeListDelete handles Asynq knowledge list delete tasks
|
||||
func (s *knowledgeService) ProcessKnowledgeListDelete(ctx context.Context, t *asynq.Task) error {
|
||||
var payload types.KnowledgeListDeletePayload
|
||||
if err := json.Unmarshal(t.Payload(), &payload); err != nil {
|
||||
logger.Errorf(ctx, "Failed to unmarshal knowledge list delete payload: %v", err)
|
||||
return err
|
||||
}
|
||||
|
||||
logger.Infof(ctx, "Processing knowledge list delete task for %d knowledge items", len(payload.KnowledgeIDs))
|
||||
|
||||
// Get tenant info
|
||||
tenant, err := s.tenantRepo.GetTenantByID(ctx, payload.TenantID)
|
||||
if err != nil {
|
||||
logger.Errorf(ctx, "Failed to get tenant %d: %v", payload.TenantID, err)
|
||||
return err
|
||||
}
|
||||
|
||||
// Set context values
|
||||
ctx = context.WithValue(ctx, types.TenantIDContextKey, payload.TenantID)
|
||||
ctx = context.WithValue(ctx, types.TenantInfoContextKey, tenant)
|
||||
|
||||
// Delete knowledge list
|
||||
if err := s.DeleteKnowledgeList(ctx, payload.KnowledgeIDs); err != nil {
|
||||
logger.Errorf(ctx, "Failed to delete knowledge list: %v", err)
|
||||
return err
|
||||
}
|
||||
|
||||
logger.Infof(ctx, "Successfully deleted %d knowledge items", len(payload.KnowledgeIDs))
|
||||
return nil
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -1,352 +0,0 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/md5"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"io"
|
||||
"mime/multipart"
|
||||
"net/http"
|
||||
"os"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
filesvc "github.com/Tencent/WeKnora/internal/application/service/file"
|
||||
"github.com/Tencent/WeKnora/internal/logger"
|
||||
"github.com/Tencent/WeKnora/internal/types"
|
||||
"github.com/Tencent/WeKnora/internal/types/interfaces"
|
||||
)
|
||||
|
||||
// isValidFileType checks if a file type is supported
|
||||
func isValidFileType(filename string) bool {
|
||||
switch strings.ToLower(getFileType(filename)) {
|
||||
case "pdf", "txt", "docx", "doc", "md", "markdown", "png", "jpg", "jpeg", "gif", "csv", "xlsx", "xls", "pptx", "ppt", "json",
|
||||
"mp3", "wav", "m4a", "flac", "ogg":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// getFileType extracts the file extension from a filename
|
||||
func getFileType(filename string) string {
|
||||
ext := strings.Split(filename, ".")
|
||||
if len(ext) < 2 {
|
||||
return "unknown"
|
||||
}
|
||||
return ext[len(ext)-1]
|
||||
}
|
||||
|
||||
// isValidURL verifies if a URL is valid
|
||||
// isValidURL 检查URL是否有效
|
||||
func isValidURL(url string) bool {
|
||||
if strings.HasPrefix(url, "http://") || strings.HasPrefix(url, "https://") {
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// calculateFileHash calculates MD5 hash of a file
|
||||
func calculateFileHash(file *multipart.FileHeader) (string, error) {
|
||||
f, err := file.Open()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer f.Close()
|
||||
|
||||
h := md5.New()
|
||||
if _, err := io.Copy(h, f); err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
// Reset file pointer for subsequent operations
|
||||
if _, err := f.Seek(0, 0); err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
return hex.EncodeToString(h.Sum(nil)), nil
|
||||
}
|
||||
|
||||
func calculateStr(strList ...string) string {
|
||||
h := md5.New()
|
||||
input := strings.Join(strList, "")
|
||||
h.Write([]byte(input))
|
||||
return hex.EncodeToString(h.Sum(nil))
|
||||
}
|
||||
|
||||
func (s *knowledgeService) getVLMConfig(ctx context.Context, kb *types.KnowledgeBase) (*types.DocParserVLMConfig, error) {
|
||||
if kb == nil {
|
||||
return nil, nil
|
||||
}
|
||||
// 兼容老版本:直接使用 ModelName 和 BaseURL
|
||||
if kb.VLMConfig.ModelName != "" && kb.VLMConfig.BaseURL != "" {
|
||||
return &types.DocParserVLMConfig{
|
||||
ModelName: kb.VLMConfig.ModelName,
|
||||
BaseURL: kb.VLMConfig.BaseURL,
|
||||
APIKey: kb.VLMConfig.APIKey,
|
||||
InterfaceType: kb.VLMConfig.InterfaceType,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// 新版本:未启用或无模型ID时返回nil
|
||||
if !kb.VLMConfig.Enabled || kb.VLMConfig.ModelID == "" {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
model, err := s.modelService.GetModelByID(ctx, kb.VLMConfig.ModelID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
interfaceType := model.Parameters.InterfaceType
|
||||
if interfaceType == "" {
|
||||
interfaceType = "openai"
|
||||
}
|
||||
|
||||
return &types.DocParserVLMConfig{
|
||||
ModelName: model.Name,
|
||||
BaseURL: model.Parameters.BaseURL,
|
||||
APIKey: model.Parameters.APIKey,
|
||||
InterfaceType: interfaceType,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *knowledgeService) buildStorageConfig(ctx context.Context, kb *types.KnowledgeBase) *types.DocParserStorageConfig {
|
||||
provider := kb.GetStorageProvider()
|
||||
if provider == "" {
|
||||
provider = "local"
|
||||
}
|
||||
|
||||
// Backward compatibility: if legacy cos_config has full params for the chosen provider, use them.
|
||||
sc := &kb.StorageConfig
|
||||
hasKBFull := false
|
||||
switch provider {
|
||||
case "cos":
|
||||
hasKBFull = sc.SecretID != "" && sc.BucketName != ""
|
||||
case "minio":
|
||||
hasKBFull = sc.BucketName != ""
|
||||
case "local":
|
||||
hasKBFull = false
|
||||
}
|
||||
|
||||
if hasKBFull {
|
||||
logger.Infof(ctx, "[storage] buildStorageConfig use legacy kb config: kb=%s provider=%s bucket=%s path_prefix=%s",
|
||||
kb.ID, provider, sc.BucketName, sc.PathPrefix)
|
||||
return &types.DocParserStorageConfig{
|
||||
Provider: strings.ToUpper(provider),
|
||||
Region: sc.Region,
|
||||
BucketName: sc.BucketName,
|
||||
AccessKeyID: sc.SecretID,
|
||||
SecretAccessKey: sc.SecretKey,
|
||||
AppID: sc.AppID,
|
||||
PathPrefix: sc.PathPrefix,
|
||||
}
|
||||
}
|
||||
|
||||
// Merge from tenant's StorageEngineConfig.
|
||||
var out types.DocParserStorageConfig
|
||||
out.Provider = strings.ToUpper(provider)
|
||||
|
||||
tenant, _ := ctx.Value(types.TenantInfoContextKey).(*types.Tenant)
|
||||
if tenant != nil && tenant.StorageEngineConfig != nil {
|
||||
sec := tenant.StorageEngineConfig
|
||||
if sec.DefaultProvider != "" && provider == "" {
|
||||
provider = strings.ToLower(strings.TrimSpace(sec.DefaultProvider))
|
||||
out.Provider = strings.ToUpper(provider)
|
||||
}
|
||||
switch provider {
|
||||
case "local":
|
||||
if sec.Local != nil {
|
||||
out.PathPrefix = sec.Local.PathPrefix
|
||||
}
|
||||
case "minio":
|
||||
if sec.MinIO != nil {
|
||||
out.BucketName = sec.MinIO.BucketName
|
||||
out.PathPrefix = sec.MinIO.PathPrefix
|
||||
if sec.MinIO.Mode == "remote" {
|
||||
out.Endpoint = sec.MinIO.Endpoint
|
||||
out.AccessKeyID = sec.MinIO.AccessKeyID
|
||||
out.SecretAccessKey = sec.MinIO.SecretAccessKey
|
||||
} else {
|
||||
out.Endpoint = os.Getenv("MINIO_ENDPOINT")
|
||||
out.AccessKeyID = os.Getenv("MINIO_ACCESS_KEY_ID")
|
||||
out.SecretAccessKey = os.Getenv("MINIO_SECRET_ACCESS_KEY")
|
||||
}
|
||||
}
|
||||
case "cos":
|
||||
if sec.COS != nil {
|
||||
out.Region = sec.COS.Region
|
||||
out.BucketName = sec.COS.BucketName
|
||||
out.AccessKeyID = sec.COS.SecretID
|
||||
out.SecretAccessKey = sec.COS.SecretKey
|
||||
out.AppID = sec.COS.AppID
|
||||
out.PathPrefix = sec.COS.PathPrefix
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
logger.Infof(ctx, "[storage] buildStorageConfig use merged tenant/global config: kb=%s provider=%s bucket=%s path_prefix=%s endpoint=%s",
|
||||
kb.ID, strings.ToLower(out.Provider), out.BucketName, out.PathPrefix, out.Endpoint)
|
||||
return &out
|
||||
}
|
||||
|
||||
// resolveFileService returns the FileService for the given knowledge base,
|
||||
// based on the KB's StorageProviderConfig (or legacy StorageConfig.Provider) and the tenant's StorageEngineConfig.
|
||||
// Falls back to the global fileSvc when no tenant-level storage config is found.
|
||||
func (s *knowledgeService) resolveFileService(ctx context.Context, kb *types.KnowledgeBase) interfaces.FileService {
|
||||
if kb == nil {
|
||||
logger.Infof(ctx, "[storage] resolveFileService fallback default: kb=nil")
|
||||
return s.fileSvc
|
||||
}
|
||||
|
||||
provider := kb.GetStorageProvider()
|
||||
|
||||
tenant, _ := ctx.Value(types.TenantInfoContextKey).(*types.Tenant)
|
||||
if provider == "" && tenant != nil && tenant.StorageEngineConfig != nil {
|
||||
provider = strings.ToLower(strings.TrimSpace(tenant.StorageEngineConfig.DefaultProvider))
|
||||
}
|
||||
|
||||
if provider == "" || tenant == nil || tenant.StorageEngineConfig == nil {
|
||||
logger.Infof(ctx, "[storage] resolveFileService fallback default: kb=%s provider=%q tenant_cfg=%v",
|
||||
kb.ID, provider, tenant != nil && tenant.StorageEngineConfig != nil)
|
||||
return s.fileSvc
|
||||
}
|
||||
|
||||
sec := tenant.StorageEngineConfig
|
||||
baseDir := strings.TrimSpace(os.Getenv("LOCAL_STORAGE_BASE_DIR"))
|
||||
svc, resolvedProvider, err := filesvc.NewFileServiceFromStorageConfig(provider, sec, baseDir)
|
||||
if err != nil {
|
||||
logger.Errorf(ctx, "Failed to create %s file service from tenant config: %v, falling back to default", provider, err)
|
||||
return s.fileSvc
|
||||
}
|
||||
logger.Infof(ctx, "[storage] resolveFileService selected: kb=%s provider=%s", kb.ID, resolvedProvider)
|
||||
return svc
|
||||
}
|
||||
|
||||
// resolveFileServiceForPath is like resolveFileService but adds a safety check:
|
||||
// if the resolved provider doesn't match what the filePath implies, fall back to
|
||||
// the provider inferred from the file path. This protects historical data when
|
||||
// tenant/KB config changes but files were stored under the old provider.
|
||||
func (s *knowledgeService) resolveFileServiceForPath(ctx context.Context, kb *types.KnowledgeBase, filePath string) interfaces.FileService {
|
||||
svc := s.resolveFileService(ctx, kb)
|
||||
if filePath == "" {
|
||||
return svc
|
||||
}
|
||||
|
||||
inferred := types.InferStorageFromFilePath(filePath)
|
||||
if inferred == "" {
|
||||
return svc
|
||||
}
|
||||
|
||||
configured := kb.GetStorageProvider()
|
||||
if configured == "" {
|
||||
tenant, _ := ctx.Value(types.TenantInfoContextKey).(*types.Tenant)
|
||||
if tenant != nil && tenant.StorageEngineConfig != nil {
|
||||
configured = strings.ToLower(strings.TrimSpace(tenant.StorageEngineConfig.DefaultProvider))
|
||||
}
|
||||
}
|
||||
if configured == "" {
|
||||
configured = strings.ToLower(strings.TrimSpace(os.Getenv("STORAGE_TYPE")))
|
||||
}
|
||||
|
||||
if configured != "" && configured != inferred {
|
||||
logger.Warnf(ctx, "[storage] FilePath format mismatch: configured=%s inferred=%s filePath=%s, using global fallback",
|
||||
configured, inferred, filePath)
|
||||
return s.fileSvc
|
||||
}
|
||||
return svc
|
||||
}
|
||||
|
||||
func IsImageType(fileType string) bool {
|
||||
switch fileType {
|
||||
case "jpg", "jpeg", "png", "gif", "webp", "bmp", "svg", "tiff":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// IsAudioType checks if a file type is an audio format
|
||||
func IsAudioType(fileType string) bool {
|
||||
switch strings.ToLower(fileType) {
|
||||
case "mp3", "wav", "m4a", "flac", "ogg":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// IsVideoType checks if a file type is a video format
|
||||
func IsVideoType(fileType string) bool {
|
||||
switch strings.ToLower(fileType) {
|
||||
case "mp4", "mov", "avi", "mkv", "webm", "wmv", "flv":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// downloadFileFromURL downloads a remote file to a temp file and returns its binary content.
|
||||
// payloadFileName and payloadFileType are in/out pointers: if they point to an empty string,
|
||||
// the function resolves the value from Content-Disposition / URL path and writes it back.
|
||||
// It does NOT perform SSRF validation — callers are responsible for that.
|
||||
func downloadFileFromURL(ctx context.Context, fileURL string, payloadFileName, payloadFileType *string) ([]byte, error) {
|
||||
httpClient := &http.Client{Timeout: 60 * time.Second}
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, fileURL, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create request for file URL: %w", err)
|
||||
}
|
||||
resp, err := httpClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to download file from URL: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("remote server returned status %d", resp.StatusCode)
|
||||
}
|
||||
|
||||
// Reject oversized files early via Content-Length
|
||||
if contentLength := resp.ContentLength; contentLength > maxFileURLSize {
|
||||
return nil, fmt.Errorf("file size %d bytes exceeds limit of %d bytes (10MB)", contentLength, maxFileURLSize)
|
||||
}
|
||||
|
||||
// Resolve fileName: payload > Content-Disposition > URL path
|
||||
if *payloadFileName == "" {
|
||||
if cd := resp.Header.Get("Content-Disposition"); cd != "" {
|
||||
*payloadFileName = extractFileNameFromContentDisposition(cd)
|
||||
}
|
||||
}
|
||||
if *payloadFileName == "" {
|
||||
*payloadFileName = extractFileNameFromURL(fileURL)
|
||||
}
|
||||
if *payloadFileType == "" && *payloadFileName != "" {
|
||||
*payloadFileType = getFileType(*payloadFileName)
|
||||
}
|
||||
|
||||
// Stream response body into a temp file, capped at maxFileURLSize
|
||||
tmpFile, err := os.CreateTemp("", "weknora-fileurl-*")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create temp file: %w", err)
|
||||
}
|
||||
tmpPath := tmpFile.Name()
|
||||
defer os.Remove(tmpPath)
|
||||
|
||||
limiter := &io.LimitedReader{R: resp.Body, N: maxFileURLSize + 1}
|
||||
written, err := io.Copy(tmpFile, limiter)
|
||||
tmpFile.Close()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to write temp file: %w", err)
|
||||
}
|
||||
if written > maxFileURLSize {
|
||||
return nil, fmt.Errorf("file size exceeds limit of 10MB")
|
||||
}
|
||||
|
||||
contentBytes, err := os.ReadFile(tmpPath)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read temp file: %w", err)
|
||||
}
|
||||
|
||||
return contentBytes, nil
|
||||
}
|
||||
Reference in New Issue
Block a user