From 3065a299352d92436e3ec18c9bc0b48350f2cfb7 Mon Sep 17 00:00:00 2001 From: Jin Hai Date: Mon, 27 Jul 2026 15:06:48 +0800 Subject: [PATCH] Go: add context, part9 (#17412) Signed-off-by: Jin Hai --- internal/dao/pipeline_operation_log.go | 19 +++--- internal/dao/search.go | 47 ++++++++------- internal/dao/search_detail_test.go | 6 +- internal/dao/skill_search_config.go | 59 ++++++++++--------- internal/dao/skill_space.go | 55 +++++++++-------- internal/handler/chat.go | 4 +- internal/handler/chat_audio.go | 4 +- internal/handler/search.go | 15 +++-- internal/handler/searchbot.go | 9 +-- internal/handler/skill_search.go | 37 +++++------- internal/ingestion/task/pipeline_executor.go | 14 +++-- .../ingestion/task/pipeline_executor_test.go | 54 ++++++++++------- .../task/pipeline_real_integration_test.go | 16 ++--- internal/service/chat.go | 34 +++++------ internal/service/chat_pipeline.go | 32 +++++----- internal/service/chat_session.go | 12 ++-- internal/service/chat_session_test.go | 2 +- internal/service/chunk/chunk.go | 20 +++---- internal/service/dataset/crud.go | 4 +- internal/service/dataset/index.go | 56 +++++++++--------- internal/service/dataset/ingestion.go | 32 +++++----- internal/service/dataset/search.go | 31 ++++++---- internal/service/dataset/update.go | 4 +- internal/service/generator.go | 4 +- internal/service/memory.go | 12 ++-- internal/service/memory_message_service.go | 2 +- internal/service/model_service.go | 32 +++++----- internal/service/model_service_test.go | 3 +- internal/service/openai_chat.go | 2 +- internal/service/related_question.go | 6 +- internal/service/search.go | 59 ++++++++++--------- internal/service/search_share_detail_test.go | 6 +- internal/service/search_test.go | 22 ++++--- internal/service/skill_indexer.go | 10 ++-- internal/service/skill_search.go | 22 +++---- internal/service/skill_space.go | 55 +++++++++-------- internal/service/user.go | 36 ----------- 37 files changed, 426 insertions(+), 411 deletions(-) diff --git a/internal/dao/pipeline_operation_log.go b/internal/dao/pipeline_operation_log.go index ead4ad893a..abffe25481 100644 --- a/internal/dao/pipeline_operation_log.go +++ b/internal/dao/pipeline_operation_log.go @@ -17,9 +17,12 @@ package dao import ( + "context" "strings" "ragflow/internal/entity" + + "gorm.io/gorm" ) // graphRaptorFakeDocID is the placeholder document_id used for dataset-level @@ -75,8 +78,8 @@ func NewPipelineOperationLogDAO() *PipelineOperationLogDAO { // GetDatasetLogsByKBID lists dataset-level (graph/raptor/mindmap) ingestion // logs for a knowledge base. Pagination is only applied when both page and // pageSize are positive, matching peewee's paginate behaviour. -func (dao *PipelineOperationLogDAO) GetDatasetLogsByKBID(kbID string, page, pageSize int, orderby string, desc bool, operationStatus []string, createDateFrom, createDateTo, keywords string) ([]*entity.PipelineOperationLog, int64, error) { - query := DB.Model(&entity.PipelineOperationLog{}). +func (dao *PipelineOperationLogDAO) GetDatasetLogsByKBID(ctx context.Context, db *gorm.DB, kbID string, page, pageSize int, orderby string, desc bool, operationStatus []string, createDateFrom, createDateTo, keywords string) ([]*entity.PipelineOperationLog, int64, error) { + query := db.WithContext(ctx).Model(&entity.PipelineOperationLog{}). Where("kb_id = ? AND document_id = ?", kbID, graphRaptorFakeDocID) if keywords != "" { @@ -115,8 +118,8 @@ func (dao *PipelineOperationLogDAO) GetDatasetLogsByKBID(kbID string, page, page } // GetFileLogsByKBID lists per-file ingestion logs for a knowledge base. -func (dao *PipelineOperationLogDAO) GetFileLogsByKBID(kbID string, page, pageSize int, orderby string, desc bool, keywords string, operationStatus []string, createDateFrom, createDateTo string) ([]*entity.PipelineOperationLog, int64, error) { - query := DB.Model(&entity.PipelineOperationLog{}). +func (dao *PipelineOperationLogDAO) GetFileLogsByKBID(ctx context.Context, db *gorm.DB, kbID string, page, pageSize int, orderby string, desc bool, keywords string, operationStatus []string, createDateFrom, createDateTo string) ([]*entity.PipelineOperationLog, int64, error) { + query := db.WithContext(ctx).Model(&entity.PipelineOperationLog{}). Where("kb_id = ?", kbID) if keywords != "" { @@ -157,15 +160,15 @@ func (dao *PipelineOperationLogDAO) GetFileLogsByKBID(kbID string, page, pageSiz } // GetByIDAndKBID fetches a single ingestion log scoped to its knowledge base. -func (dao *PipelineOperationLogDAO) GetByIDAndKBID(logID, kbID string) (*entity.PipelineOperationLog, error) { +func (dao *PipelineOperationLogDAO) GetByIDAndKBID(ctx context.Context, db *gorm.DB, logID, kbID string) (*entity.PipelineOperationLog, error) { var log entity.PipelineOperationLog - if err := DB.Where("id = ? AND kb_id = ?", logID, kbID).First(&log).Error; err != nil { + if err := db.WithContext(ctx).Where("id = ? AND kb_id = ?", logID, kbID).First(&log).Error; err != nil { return nil, err } return &log, nil } // Create inserts a new pipeline operation log. -func (dao *PipelineOperationLogDAO) Create(log *entity.PipelineOperationLog) error { - return DB.Create(log).Error +func (dao *PipelineOperationLogDAO) Create(ctx context.Context, db *gorm.DB, log *entity.PipelineOperationLog) error { + return db.WithContext(ctx).Create(log).Error } diff --git a/internal/dao/search.go b/internal/dao/search.go index 46f1e241f2..e12aa99707 100644 --- a/internal/dao/search.go +++ b/internal/dao/search.go @@ -17,8 +17,11 @@ package dao import ( + "context" "ragflow/internal/entity" "strings" + + "gorm.io/gorm" ) // SearchDAO search data access object @@ -45,12 +48,12 @@ type SearchDetailRow struct { } // ListByTenantIDs list searches by tenant IDs with pagination and filtering -func (dao *SearchDAO) ListByTenantIDs(tenantIDs []string, userID string, page, pageSize int, orderby string, desc bool, keywords string) ([]*entity.SearchListItem, int64, error) { +func (dao *SearchDAO) ListByTenantIDs(ctx context.Context, db *gorm.DB, tenantIDs []string, userID string, page, pageSize int, orderby string, desc bool, keywords string) ([]*entity.SearchListItem, int64, error) { var searches []*entity.SearchListItem var total int64 // Build query with join to user table for nickname and avatar - query := DB.Model(&entity.Search{}). + query := db.WithContext(ctx).Model(&entity.Search{}). Select(` search.*, user.nickname, @@ -97,11 +100,11 @@ func (dao *SearchDAO) ListByTenantIDs(tenantIDs []string, userID string, page, p } // ListByOwnerIDs list searches by owner IDs with filtering (manual pagination) -func (dao *SearchDAO) ListByOwnerIDs(ownerIDs []string, userID string, orderby string, desc bool, keywords string) ([]*entity.SearchListItem, int64, error) { +func (dao *SearchDAO) ListByOwnerIDs(ctx context.Context, db *gorm.DB, ownerIDs []string, userID string, orderby string, desc bool, keywords string) ([]*entity.SearchListItem, int64, error) { var searches []*entity.SearchListItem // Build query with join to user table - query := DB.Model(&entity.Search{}). + query := db.WithContext(ctx).Model(&entity.Search{}). Select(` search.*, user.nickname, @@ -136,9 +139,9 @@ func (dao *SearchDAO) ListByOwnerIDs(ownerIDs []string, userID string, orderby s } // GetByID gets search by ID -func (dao *SearchDAO) GetByID(id string) (*entity.Search, error) { +func (dao *SearchDAO) GetByID(ctx context.Context, db *gorm.DB, id string) (*entity.Search, error) { var search entity.Search - err := DB.Where("id = ?", id).First(&search).Error + err := db.WithContext(ctx).Where("id = ?", id).First(&search).Error if err != nil { return nil, err } @@ -147,9 +150,9 @@ func (dao *SearchDAO) GetByID(id string) (*entity.Search, error) { // GetDetailByID retrieves the share-detail payload by joining the search app // with its owner profile, matching Python SearchService.get_detail. -func (dao *SearchDAO) GetDetailByID(searchID string) (*SearchDetailRow, error) { +func (dao *SearchDAO) GetDetailByID(ctx context.Context, db *gorm.DB, searchID string) (*SearchDetailRow, error) { var detail SearchDetailRow - err := DB.Table("search"). + err := db.WithContext(ctx).Table("search"). Select(` search.id, search.avatar, @@ -175,30 +178,30 @@ func (dao *SearchDAO) GetDetailByID(searchID string) (*SearchDetailRow, error) { } // GetByNameAndTenant gets search by name and tenant ID -func (dao *SearchDAO) GetByNameAndTenant(name string, tenantID string) ([]*entity.Search, error) { +func (dao *SearchDAO) GetByNameAndTenant(ctx context.Context, db *gorm.DB, name string, tenantID string) ([]*entity.Search, error) { var searches []*entity.Search - err := DB.Where("name = ? AND tenant_id = ? AND status = ?", name, tenantID, "1").Find(&searches).Error + err := db.WithContext(ctx).Where("name = ? AND tenant_id = ? AND status = ?", name, tenantID, "1").Find(&searches).Error return searches, err } // Create creates a new search -func (dao *SearchDAO) Create(search *entity.Search) error { - return DB.Create(search).Error +func (dao *SearchDAO) Create(ctx context.Context, db *gorm.DB, search *entity.Search) error { + return db.WithContext(ctx).Create(search).Error } // QueryByTenantIDAndID checks if a search exists with given tenant_id and id // Reference: Python SearchService.query(tenant_id=tenant.tenant_id, id=search_id) // Used for permission verification in detail API -func (dao *SearchDAO) QueryByTenantIDAndID(tenantID string, searchID string) ([]*entity.Search, error) { +func (dao *SearchDAO) QueryByTenantIDAndID(ctx context.Context, db *gorm.DB, tenantID string, searchID string) ([]*entity.Search, error) { var searches []*entity.Search - err := DB.Where("tenant_id = ? AND id = ? AND status = ?", tenantID, searchID, "1").Find(&searches).Error + err := db.WithContext(ctx).Where("tenant_id = ? AND id = ? AND status = ?", tenantID, searchID, "1").Find(&searches).Error return searches, err } // DeleteByID deletes a search by ID (soft delete by setting status to "0") // Reference: Python common_service.py::delete_by_id -func (dao *SearchDAO) DeleteByID(tenantID, id string) error { - return DB.Model(&entity.Search{}).Where("tenant_id = ? AND id = ?", tenantID, id).Update("status", "0").Error +func (dao *SearchDAO) DeleteByID(ctx context.Context, db *gorm.DB, tenantID, id string) error { + return db.WithContext(ctx).Model(&entity.Search{}).Where("tenant_id = ? AND id = ?", tenantID, id).Update("status", "0").Error } // Accessible4Deletion checks if a search can be deleted by a specific user @@ -206,9 +209,9 @@ func (dao *SearchDAO) DeleteByID(tenantID, id string) error { // Returns true if the search exists, is valid, and was created by the user. // A missing or non-owned search returns (false, nil) so callers can distinguish // "not authorized" from a genuine database error (which is returned as the error). -func (dao *SearchDAO) Accessible4Deletion(searchID string, userID string) (bool, error) { +func (dao *SearchDAO) Accessible4Deletion(ctx context.Context, db *gorm.DB, searchID string, userID string) (bool, error) { var count int64 - err := DB.Model(&entity.Search{}). + err := db.WithContext(ctx).Model(&entity.Search{}). Where("id = ? AND created_by = ? AND status = ?", searchID, userID, "1"). Count(&count).Error if err != nil { @@ -219,9 +222,9 @@ func (dao *SearchDAO) Accessible4Deletion(searchID string, userID string) (bool, // GetByTenantIDAndID gets search by tenant ID and search ID // Reference: Python SearchService.query(tenant_id=tenant_id, id=search_id) -func (dao *SearchDAO) GetByTenantIDAndID(tenantID string, searchID string) (*entity.Search, error) { +func (dao *SearchDAO) GetByTenantIDAndID(ctx context.Context, db *gorm.DB, tenantID string, searchID string) (*entity.Search, error) { var search entity.Search - err := DB.Where("tenant_id = ? AND id = ? AND status = ?", tenantID, searchID, "1").First(&search).Error + err := db.WithContext(ctx).Where("tenant_id = ? AND id = ? AND status = ?", tenantID, searchID, "1").First(&search).Error if err != nil { return nil, err } @@ -230,6 +233,6 @@ func (dao *SearchDAO) GetByTenantIDAndID(tenantID string, searchID string) (*ent // UpdateByID updates search by ID // Reference: Python common_service.py::update_by_id -func (dao *SearchDAO) UpdateByID(id string, updates map[string]interface{}) error { - return DB.Model(&entity.Search{}).Where("id = ?", id).Updates(updates).Error +func (dao *SearchDAO) UpdateByID(ctx context.Context, db *gorm.DB, id string, updates map[string]interface{}) error { + return db.WithContext(ctx).Model(&entity.Search{}).Where("id = ?", id).Updates(updates).Error } diff --git a/internal/dao/search_detail_test.go b/internal/dao/search_detail_test.go index 04d81b678b..e08b361ae3 100644 --- a/internal/dao/search_detail_test.go +++ b/internal/dao/search_detail_test.go @@ -75,7 +75,8 @@ func TestSearchDAOGetDetailByID(t *testing.T) { t.Fatalf("failed to create search: %v", err) } - detail, err := NewSearchDAO().GetDetailByID("search-1") + ctx := t.Context() + detail, err := NewSearchDAO().GetDetailByID(ctx, db, "search-1") if err != nil { t.Fatalf("GetDetailByID failed: %v", err) } @@ -128,7 +129,8 @@ func TestSearchDAOGetDetailByIDReturnsNilWhenJoinedUserIsInactive(t *testing.T) t.Fatalf("failed to create search: %v", err) } - detail, err := NewSearchDAO().GetDetailByID("search-1") + ctx := t.Context() + detail, err := NewSearchDAO().GetDetailByID(ctx, db, "search-1") if err != nil { t.Fatalf("GetDetailByID failed: %v", err) } diff --git a/internal/dao/skill_search_config.go b/internal/dao/skill_search_config.go index 8a3939de1a..b4db9d451c 100644 --- a/internal/dao/skill_search_config.go +++ b/internal/dao/skill_search_config.go @@ -17,9 +17,12 @@ package dao import ( + "context" "ragflow/internal/entity" "ragflow/internal/utility" "strings" + + "gorm.io/gorm" ) // SkillSearchConfigDAO data access object for skill search config @@ -41,14 +44,14 @@ func NewSkillSearchConfigDAO() *SkillSearchConfigDAO { } // Create creates a new skill search config -func (dao *SkillSearchConfigDAO) Create(config *entity.SkillSearchConfig) error { - return DB.Create(config).Error +func (dao *SkillSearchConfigDAO) Create(ctx context.Context, db *gorm.DB, config *entity.SkillSearchConfig) error { + return db.WithContext(ctx).Create(config).Error } // GetByID retrieves a skill search config by ID -func (dao *SkillSearchConfigDAO) GetByID(id string) (*entity.SkillSearchConfig, error) { +func (dao *SkillSearchConfigDAO) GetByID(ctx context.Context, db *gorm.DB, id string) (*entity.SkillSearchConfig, error) { var config entity.SkillSearchConfig - err := DB.Where("id = ? AND status = ?", id, "1").First(&config).Error + err := db.WithContext(ctx).Where("id = ? AND status = ?", id, "1").First(&config).Error if err != nil { return nil, err } @@ -56,9 +59,9 @@ func (dao *SkillSearchConfigDAO) GetByID(id string) (*entity.SkillSearchConfig, } // GetByTenantID retrieves a skill search config by tenant ID -func (dao *SkillSearchConfigDAO) GetByTenantID(tenantID, spaceID string) (*entity.SkillSearchConfig, error) { +func (dao *SkillSearchConfigDAO) GetByTenantID(ctx context.Context, db *gorm.DB, tenantID, spaceID string) (*entity.SkillSearchConfig, error) { var config entity.SkillSearchConfig - err := DB.Where("tenant_id = ? AND space_id = ? AND status = ?", tenantID, normalizeSpaceID(spaceID), "1").First(&config).Error + err := db.WithContext(ctx).Where("tenant_id = ? AND space_id = ? AND status = ?", tenantID, normalizeSpaceID(spaceID), "1").First(&config).Error if err != nil { return nil, err } @@ -67,15 +70,15 @@ func (dao *SkillSearchConfigDAO) GetByTenantID(tenantID, spaceID string) (*entit // GetLatestByTenantID retrieves the latest skill search config by tenant ID (ordered by update_time desc) // Prioritizes configs with non-empty embd_id to return user-saved configs over auto-created ones -func (dao *SkillSearchConfigDAO) GetLatestByTenantID(tenantID, spaceID string) (*entity.SkillSearchConfig, error) { +func (dao *SkillSearchConfigDAO) GetLatestByTenantID(ctx context.Context, db *gorm.DB, tenantID, spaceID string) (*entity.SkillSearchConfig, error) { var config entity.SkillSearchConfig // First try to get the latest config with non-empty embd_id (user-saved config) - err := DB.Where("tenant_id = ? AND space_id = ? AND status = ? AND embd_id != ?", tenantID, normalizeSpaceID(spaceID), "1", "").Order("update_time desc").First(&config).Error + err := db.WithContext(ctx).Where("tenant_id = ? AND space_id = ? AND status = ? AND embd_id != ?", tenantID, normalizeSpaceID(spaceID), "1", "").Order("update_time desc").First(&config).Error if err == nil { return &config, nil } // If no user-saved config found, get any config - err = DB.Where("tenant_id = ? AND space_id = ? AND status = ?", tenantID, normalizeSpaceID(spaceID), "1").Order("update_time desc").First(&config).Error + err = db.WithContext(ctx).Where("tenant_id = ? AND space_id = ? AND status = ?", tenantID, normalizeSpaceID(spaceID), "1").Order("update_time desc").First(&config).Error if err != nil { return nil, err } @@ -83,9 +86,9 @@ func (dao *SkillSearchConfigDAO) GetLatestByTenantID(tenantID, spaceID string) ( } // GetByTenantAndEmbdID retrieves a skill search config by tenant ID and embedding ID -func (dao *SkillSearchConfigDAO) GetByTenantAndEmbdID(tenantID, spaceID, embdID string) (*entity.SkillSearchConfig, error) { +func (dao *SkillSearchConfigDAO) GetByTenantAndEmbdID(ctx context.Context, db *gorm.DB, tenantID, spaceID, embdID string) (*entity.SkillSearchConfig, error) { var config entity.SkillSearchConfig - err := DB.Where("tenant_id = ? AND space_id = ? AND embd_id = ? AND status = ?", tenantID, normalizeSpaceID(spaceID), embdID, "1").First(&config).Error + err := db.WithContext(ctx).Where("tenant_id = ? AND space_id = ? AND embd_id = ? AND status = ?", tenantID, normalizeSpaceID(spaceID), embdID, "1").First(&config).Error if err != nil { return nil, err } @@ -93,19 +96,19 @@ func (dao *SkillSearchConfigDAO) GetByTenantAndEmbdID(tenantID, spaceID, embdID } // GetOrCreate retrieves existing config or creates default one -func (dao *SkillSearchConfigDAO) GetOrCreate(tenantID, spaceID, embdID string) (*entity.SkillSearchConfig, error) { +func (dao *SkillSearchConfigDAO) GetOrCreate(ctx context.Context, db *gorm.DB, tenantID, spaceID, embdID string) (*entity.SkillSearchConfig, error) { spaceID = normalizeSpaceID(spaceID) - config, err := dao.GetByTenantAndEmbdID(tenantID, spaceID, embdID) + config, err := dao.GetByTenantAndEmbdID(ctx, db, tenantID, spaceID, embdID) if err == nil { return config, nil } // Create default config - return dao.CreateWithTenantSpace(tenantID, spaceID, embdID) + return dao.CreateWithTenantSpace(ctx, db, tenantID, spaceID, embdID) } // CreateWithTenantSpace creates a new config for tenant+space -func (dao *SkillSearchConfigDAO) CreateWithTenantSpace(tenantID, spaceID, embdID string) (*entity.SkillSearchConfig, error) { +func (dao *SkillSearchConfigDAO) CreateWithTenantSpace(ctx context.Context, db *gorm.DB, tenantID, spaceID, embdID string) (*entity.SkillSearchConfig, error) { spaceID = normalizeSpaceID(spaceID) defaultFieldConfig := entity.DefaultFieldConfig() fieldConfigMap := entity.JSONMap{ @@ -139,46 +142,46 @@ func (dao *SkillSearchConfigDAO) CreateWithTenantSpace(tenantID, spaceID, embdID Status: "1", } - if err := dao.Create(defaultConfig); err != nil { + if err := dao.Create(ctx, db, defaultConfig); err != nil { return nil, err } return defaultConfig, nil } // DeleteAllByTenantSpace deletes all configs for a tenant+space (for cleanup before creating new one) -func (dao *SkillSearchConfigDAO) DeleteAllByTenantSpace(tenantID, spaceID string) error { +func (dao *SkillSearchConfigDAO) DeleteAllByTenantSpace(ctx context.Context, db *gorm.DB, tenantID, spaceID string) error { spaceID = normalizeSpaceID(spaceID) - return DB.Model(&entity.SkillSearchConfig{}). + return db.WithContext(ctx).Model(&entity.SkillSearchConfig{}). Where("tenant_id = ? AND space_id = ?", tenantID, spaceID). Update("status", "0").Error } // DeleteAllByTenantSpaceExceptID deletes all active configs for a tenant+space except the specified ID -func (dao *SkillSearchConfigDAO) DeleteAllByTenantSpaceExceptID(tenantID, spaceID, exceptID string) error { +func (dao *SkillSearchConfigDAO) DeleteAllByTenantSpaceExceptID(ctx context.Context, db *gorm.DB, tenantID, spaceID, exceptID string) error { spaceID = normalizeSpaceID(spaceID) - return DB.Model(&entity.SkillSearchConfig{}). + return db.WithContext(ctx).Model(&entity.SkillSearchConfig{}). Where("tenant_id = ? AND space_id = ? AND id != ? AND status = ?", tenantID, spaceID, exceptID, "1"). Update("status", "0").Error } // Update updates a skill search config with the given updates map -func (dao *SkillSearchConfigDAO) Update(id string, updates map[string]interface{}) error { - return DB.Model(&entity.SkillSearchConfig{}).Where("id = ? AND status = ?", id, "1").Updates(updates).Error +func (dao *SkillSearchConfigDAO) Update(ctx context.Context, db *gorm.DB, id string, updates map[string]interface{}) error { + return db.WithContext(ctx).Model(&entity.SkillSearchConfig{}).Where("id = ? AND status = ?", id, "1").Updates(updates).Error } // UpdateByTenantID updates config by tenant ID -func (dao *SkillSearchConfigDAO) UpdateByTenantID(tenantID, spaceID string, updates map[string]interface{}) error { - result := DB.Model(&entity.SkillSearchConfig{}).Where("tenant_id = ? AND space_id = ? AND status = ?", tenantID, normalizeSpaceID(spaceID), "1").Updates(updates) +func (dao *SkillSearchConfigDAO) UpdateByTenantID(ctx context.Context, db *gorm.DB, tenantID, spaceID string, updates map[string]interface{}) error { + result := db.WithContext(ctx).Model(&entity.SkillSearchConfig{}).Where("tenant_id = ? AND space_id = ? AND status = ?", tenantID, normalizeSpaceID(spaceID), "1").Updates(updates) return result.Error } // UpdateByTenantAndEmbdID updates config by tenant ID and embedding ID -func (dao *SkillSearchConfigDAO) UpdateByTenantAndEmbdID(tenantID, spaceID, embdID string, updates map[string]interface{}) error { - result := DB.Model(&entity.SkillSearchConfig{}).Where("tenant_id = ? AND space_id = ? AND embd_id = ? AND status = ?", tenantID, normalizeSpaceID(spaceID), embdID, "1").Updates(updates) +func (dao *SkillSearchConfigDAO) UpdateByTenantAndEmbdID(ctx context.Context, db *gorm.DB, tenantID, spaceID, embdID string, updates map[string]interface{}) error { + result := db.WithContext(ctx).Model(&entity.SkillSearchConfig{}).Where("tenant_id = ? AND space_id = ? AND embd_id = ? AND status = ?", tenantID, normalizeSpaceID(spaceID), embdID, "1").Updates(updates) return result.Error } // Delete deletes a skill search config by ID (soft delete) -func (dao *SkillSearchConfigDAO) Delete(id string) error { - return DB.Model(&entity.SkillSearchConfig{}).Where("id = ?", id).Update("status", "0").Error +func (dao *SkillSearchConfigDAO) Delete(ctx context.Context, db *gorm.DB, id string) error { + return db.WithContext(ctx).Model(&entity.SkillSearchConfig{}).Where("id = ?", id).Update("status", "0").Error } diff --git a/internal/dao/skill_space.go b/internal/dao/skill_space.go index 937dfdb94b..c3637c073e 100644 --- a/internal/dao/skill_space.go +++ b/internal/dao/skill_space.go @@ -17,7 +17,10 @@ package dao import ( + "context" "ragflow/internal/entity" + + "gorm.io/gorm" ) // SkillSpaceDAO data access object for skills space @@ -29,14 +32,14 @@ func NewSkillSpaceDAO() *SkillSpaceDAO { } // Create creates a new skills space -func (dao *SkillSpaceDAO) Create(space *entity.SkillSpace) error { - return DB.Create(space).Error +func (dao *SkillSpaceDAO) Create(ctx context.Context, db *gorm.DB, space *entity.SkillSpace) error { + return db.WithContext(ctx).Create(space).Error } // GetByID retrieves a skills space by ID (active only) -func (dao *SkillSpaceDAO) GetByID(id string) (*entity.SkillSpace, error) { +func (dao *SkillSpaceDAO) GetByID(ctx context.Context, db *gorm.DB, id string) (*entity.SkillSpace, error) { var space entity.SkillSpace - err := DB.Where("id = ? AND status = ?", id, entity.SpaceStatusActive).First(&space).Error + err := db.WithContext(ctx).Where("id = ? AND status = ?", id, entity.SpaceStatusActive).First(&space).Error if err != nil { return nil, err } @@ -44,16 +47,16 @@ func (dao *SkillSpaceDAO) GetByID(id string) (*entity.SkillSpace, error) { } // GetByTenantID retrieves all skills spaces by tenant ID (active only) -func (dao *SkillSpaceDAO) GetByTenantID(tenantID string) ([]*entity.SkillSpace, error) { +func (dao *SkillSpaceDAO) GetByTenantID(ctx context.Context, db *gorm.DB, tenantID string) ([]*entity.SkillSpace, error) { var spaces []*entity.SkillSpace - err := DB.Where("tenant_id = ? AND status = ?", tenantID, entity.SpaceStatusActive).Order("create_time DESC").Find(&spaces).Error + err := db.WithContext(ctx).Where("tenant_id = ? AND status = ?", tenantID, entity.SpaceStatusActive).Order("create_time DESC").Find(&spaces).Error return spaces, err } // GetByTenantAndName retrieves a skills space by tenant ID and name (active only) -func (dao *SkillSpaceDAO) GetByTenantAndName(tenantID, name string) (*entity.SkillSpace, error) { +func (dao *SkillSpaceDAO) GetByTenantAndName(ctx context.Context, db *gorm.DB, tenantID, name string) (*entity.SkillSpace, error) { var space entity.SkillSpace - err := DB.Where("tenant_id = ? AND name = ? AND status = ?", tenantID, name, entity.SpaceStatusActive).First(&space).Error + err := db.WithContext(ctx).Where("tenant_id = ? AND name = ? AND status = ?", tenantID, name, entity.SpaceStatusActive).First(&space).Error if err != nil { return nil, err } @@ -61,9 +64,9 @@ func (dao *SkillSpaceDAO) GetByTenantAndName(tenantID, name string) (*entity.Ski } // GetByTenantAndNameAnyStatus retrieves a skills space by tenant ID and name regardless of status -func (dao *SkillSpaceDAO) GetByTenantAndNameAnyStatus(tenantID, name string) (*entity.SkillSpace, error) { +func (dao *SkillSpaceDAO) GetByTenantAndNameAnyStatus(ctx context.Context, db *gorm.DB, tenantID, name string) (*entity.SkillSpace, error) { var space entity.SkillSpace - err := DB.Where("tenant_id = ? AND name = ?", tenantID, name).First(&space).Error + err := db.WithContext(ctx).Where("tenant_id = ? AND name = ?", tenantID, name).First(&space).Error if err != nil { return nil, err } @@ -71,9 +74,9 @@ func (dao *SkillSpaceDAO) GetByTenantAndNameAnyStatus(tenantID, name string) (*e } // GetByIDAnyStatus retrieves a skills space by ID regardless of status -func (dao *SkillSpaceDAO) GetByIDAnyStatus(id string) (*entity.SkillSpace, error) { +func (dao *SkillSpaceDAO) GetByIDAnyStatus(ctx context.Context, db *gorm.DB, id string) (*entity.SkillSpace, error) { var space entity.SkillSpace - err := DB.Where("id = ?", id).First(&space).Error + err := db.WithContext(ctx).Where("id = ?", id).First(&space).Error if err != nil { return nil, err } @@ -81,9 +84,9 @@ func (dao *SkillSpaceDAO) GetByIDAnyStatus(id string) (*entity.SkillSpace, error } // GetByFolderID retrieves a skills space by folder ID (active only) -func (dao *SkillSpaceDAO) GetByFolderID(folderID string) (*entity.SkillSpace, error) { +func (dao *SkillSpaceDAO) GetByFolderID(ctx context.Context, db *gorm.DB, folderID string) (*entity.SkillSpace, error) { var space entity.SkillSpace - err := DB.Where("folder_id = ? AND status = ?", folderID, entity.SpaceStatusActive).First(&space).Error + err := db.WithContext(ctx).Where("folder_id = ? AND status = ?", folderID, entity.SpaceStatusActive).First(&space).Error if err != nil { return nil, err } @@ -91,24 +94,24 @@ func (dao *SkillSpaceDAO) GetByFolderID(folderID string) (*entity.SkillSpace, er } // Update updates a skills space -func (dao *SkillSpaceDAO) Update(space *entity.SkillSpace) error { - return DB.Save(space).Error +func (dao *SkillSpaceDAO) Update(ctx context.Context, db *gorm.DB, space *entity.SkillSpace) error { + return db.WithContext(ctx).Save(space).Error } // UpdateByID updates skills space by ID -func (dao *SkillSpaceDAO) UpdateByID(id string, updates map[string]interface{}) error { - return DB.Model(&entity.SkillSpace{}).Where("id = ?", id).Updates(updates).Error +func (dao *SkillSpaceDAO) UpdateByID(ctx context.Context, db *gorm.DB, id string, updates map[string]interface{}) error { + return db.WithContext(ctx).Model(&entity.SkillSpace{}).Where("id = ?", id).Updates(updates).Error } // Delete deletes a skills space by ID (soft delete) -func (dao *SkillSpaceDAO) Delete(id string) error { - return DB.Model(&entity.SkillSpace{}).Where("id = ?", id).Update("status", entity.SpaceStatusDeleted).Error +func (dao *SkillSpaceDAO) Delete(ctx context.Context, db *gorm.DB, id string) error { + return db.WithContext(ctx).Model(&entity.SkillSpace{}).Where("id = ?", id).Update("status", entity.SpaceStatusDeleted).Error } // CASStatus performs a compare-and-swap on the space status atomically // Returns true if the update was applied, false if the current status didn't match expected -func (dao *SkillSpaceDAO) CASStatus(id string, expectedStatus, newStatus string) (bool, error) { - result := DB.Model(&entity.SkillSpace{}). +func (dao *SkillSpaceDAO) CASStatus(ctx context.Context, db *gorm.DB, id string, expectedStatus, newStatus string) (bool, error) { + result := db.WithContext(ctx).Model(&entity.SkillSpace{}). Where("id = ? AND status = ?", id, expectedStatus). Update("status", newStatus) if result.Error != nil { @@ -119,13 +122,13 @@ func (dao *SkillSpaceDAO) CASStatus(id string, expectedStatus, newStatus string) // DeletePermanentByName permanently deletes a skills space by tenant ID and name // This is used to clean up previously deleted spaces (only deletes status='0' deleted spaces, NOT deleting spaces) -func (dao *SkillSpaceDAO) DeletePermanentByName(tenantID, name string) error { - return DB.Unscoped().Where("tenant_id = ? AND name = ? AND status = ?", tenantID, name, entity.SpaceStatusDeleted).Delete(&entity.SkillSpace{}).Error +func (dao *SkillSpaceDAO) DeletePermanentByName(ctx context.Context, db *gorm.DB, tenantID, name string) error { + return db.WithContext(ctx).Unscoped().Where("tenant_id = ? AND name = ? AND status = ?", tenantID, name, entity.SpaceStatusDeleted).Delete(&entity.SkillSpace{}).Error } // CountByTenant counts skills spaces by tenant ID -func (dao *SkillSpaceDAO) CountByTenant(tenantID string) (int64, error) { +func (dao *SkillSpaceDAO) CountByTenant(ctx context.Context, db *gorm.DB, tenantID string) (int64, error) { var count int64 - err := DB.Model(&entity.SkillSpace{}).Where("tenant_id = ? AND status = ?", tenantID, entity.SpaceStatusActive).Count(&count).Error + err := db.WithContext(ctx).Model(&entity.SkillSpace{}).Where("tenant_id = ? AND status = ?", tenantID, entity.SpaceStatusActive).Count(&count).Error return count, err } diff --git a/internal/handler/chat.go b/internal/handler/chat.go index aefe305718..ab3938566c 100644 --- a/internal/handler/chat.go +++ b/internal/handler/chat.go @@ -178,6 +178,7 @@ func (h *ChatHandler) MindMap(c *gin.Context) { return } + ctx := c.Request.Context() searchConfig := map[string]interface{}{} modelTenantID := user.ID if req.SearchID != "" { @@ -185,7 +186,7 @@ func (h *ChatHandler) MindMap(c *gin.Context) { jsonInternalError(c, fmt.Errorf("search service not configured")) return } - detail, err := h.searchSvc.GetDetail(req.SearchID) + detail, err := h.searchSvc.GetDetail(ctx, req.SearchID) if err != nil { jsonInternalError(c, err) return @@ -202,7 +203,6 @@ func (h *ChatHandler) MindMap(c *gin.Context) { return } - ctx := c.Request.Context() mindMap, err := runMindMap(ctx, mindMapRunConfig{ Question: req.Question, KbIDs: kbIDs, diff --git a/internal/handler/chat_audio.go b/internal/handler/chat_audio.go index 2197e51a5f..c2090734fd 100644 --- a/internal/handler/chat_audio.go +++ b/internal/handler/chat_audio.go @@ -72,7 +72,7 @@ func (h *ChatHandler) ChatAudioSpeech(c *gin.Context) { return } - driver, modelName, apiConfig, _, err := h.llm.GetTenantDefaultModelByType(user.ID, entity.ModelTypeTTS) + driver, modelName, apiConfig, _, err := h.llm.GetTenantDefaultModelByType(ctx, user.ID, entity.ModelTypeTTS) if err != nil { common.ErrorWithCode(c, common.CodeDataError, err.Error()) return @@ -201,7 +201,7 @@ func (h *ChatHandler) ChatAudioTranscription(c *gin.Context) { return } - driver, modelName, apiConfig, _, err := h.llm.GetTenantDefaultModelByType(user.ID, entity.ModelTypeSpeech2Text) + driver, modelName, apiConfig, _, err := h.llm.GetTenantDefaultModelByType(ctx, user.ID, entity.ModelTypeSpeech2Text) if err != nil { common.ErrorWithCode(c, common.CodeDataError, err.Error()) return diff --git a/internal/handler/search.go b/internal/handler/search.go index 9616e48b9d..5477c54d41 100644 --- a/internal/handler/search.go +++ b/internal/handler/search.go @@ -128,7 +128,8 @@ func (h *SearchHandler) ListSearches(c *gin.Context) { } // List search apps with filtering - result, err := h.searchService.ListSearches(userID, keywords, page, pageSize, orderby, desc, ownerIDs) + ctx := c.Request.Context() + result, err := h.searchService.ListSearches(ctx, userID, keywords, page, pageSize, orderby, desc, ownerIDs) if err != nil { common.ResponseWithHttpCodeData(c, http.StatusInternalServerError, 500, nil, err.Error()) return @@ -174,7 +175,8 @@ func (h *SearchHandler) CreateSearch(c *gin.Context) { } // Create search (same as Python SearchService.save within DB.atomic()) - result, err := h.searchService.CreateSearch(userID, req.Name, req.Description) + ctx := c.Request.Context() + result, err := h.searchService.CreateSearch(ctx, userID, req.Name, req.Description) if err != nil { common.ResponseWithHttpCodeData(c, http.StatusInternalServerError, common.CodeBadRequest, nil, err.Error()) return @@ -210,7 +212,8 @@ func (h *SearchHandler) GetSearch(c *gin.Context) { } // Get search detail with permission check - search, err := h.searchService.GetSearchDetail(userID, searchID) + ctx := c.Request.Context() + search, err := h.searchService.GetSearchDetail(ctx, userID, searchID) if err != nil { // Check if it's a permission error if err.Error() == "has no permission for this operation" { @@ -268,7 +271,8 @@ func (h *SearchHandler) DeleteSearch(c *gin.Context) { } // Delete search with permission check - err := h.searchService.DeleteSearch(userID, searchID) + ctx := c.Request.Context() + err := h.searchService.DeleteSearch(ctx, userID, searchID) if err != nil { // Check if it's an authorization error if err.Error() == "no authorization" { @@ -324,7 +328,8 @@ func (h *SearchHandler) UpdateSearch(c *gin.Context) { } // Update search - updatedSearch, err := h.searchService.UpdateSearch(userID, searchID, &req) + ctx := c.Request.Context() + updatedSearch, err := h.searchService.UpdateSearch(ctx, userID, searchID, &req) if err != nil { errMsg := err.Error() switch errMsg { diff --git a/internal/handler/searchbot.go b/internal/handler/searchbot.go index db814b14e6..bc0b95d41c 100644 --- a/internal/handler/searchbot.go +++ b/internal/handler/searchbot.go @@ -250,7 +250,8 @@ func (h *SearchBotHandler) Ask(c *gin.Context) { // Resolve chat model ID. modelID := "" if req.SearchID != "" && h.searchSvc != nil { - if detail, err := h.searchSvc.GetDetail(req.SearchID); err == nil { + ctx := c.Request.Context() + if detail, err := h.searchSvc.GetDetail(ctx, req.SearchID); err == nil { if sc, ok := detail["search_config"].(map[string]interface{}); ok { if cid, ok := sc["chat_id"].(string); ok && cid != "" { modelID = cid @@ -350,7 +351,8 @@ func (h *SearchBotHandler) MindMap(c *gin.Context) { jsonInternalError(c, fmt.Errorf("search service not configured")) return } - detail, err := h.searchSvc.GetDetail(req.SearchID) + ctx := c.Request.Context() + detail, err := h.searchSvc.GetDetail(ctx, req.SearchID) if err != nil { jsonInternalError(c, err) return @@ -395,8 +397,7 @@ func (h *SearchBotHandler) SearchBotDetail(c *gin.Context) { common.ResponseWithCodeData(c, code, nil, "Authentication error: API key is invalid!") return } - - detail, err := h.searchSvc.GetSearchShareDetail(user.ID, searchID) + detail, err := h.searchSvc.GetSearchShareDetail(ctx, user.ID, searchID) if err != nil { switch err.Error() { case "has no permission for this operation": diff --git a/internal/handler/skill_search.go b/internal/handler/skill_search.go index 9ef2a3a2d0..51e5196fb6 100644 --- a/internal/handler/skill_search.go +++ b/internal/handler/skill_search.go @@ -51,8 +51,6 @@ func NewSkillSearchHandler(docEngine engine.DocEngine, spaceRemover file.DocRemo // @Summary Get Skill Search Config // @Description Get the search configuration for skills // @Tags skill-search -// @Accept json -// @Produce json // @Security ApiKeyAuth // @Param embd_id query string true "Embedding Model ID" // @Param space_id query string false "Skill Space ID" @@ -68,7 +66,8 @@ func (h *SkillSearchHandler) GetConfig(c *gin.Context) { embdID := c.Query("embd_id") spaceID := c.Query("space_id") - result, code, err := h.searchService.GetConfig(user.ID, spaceID, embdID) + ctx := c.Request.Context() + result, code, err := h.searchService.GetConfig(ctx, user.ID, spaceID, embdID) if err != nil { common.ErrorWithCode(c, code, err.Error()) return @@ -81,8 +80,6 @@ func (h *SkillSearchHandler) GetConfig(c *gin.Context) { // @Summary Update Skill Search Config // @Description Update the search configuration for skills // @Tags skill-search -// @Accept json -// @Produce json // @Security ApiKeyAuth // @Param request body service.UpdateConfigRequest true "config info" // @Success 200 {object} map[string]interface{} @@ -102,7 +99,8 @@ func (h *SkillSearchHandler) UpdateConfig(c *gin.Context) { req.TenantID = user.ID - result, code, err := h.searchService.UpdateConfig(&req) + ctx := c.Request.Context() + result, code, err := h.searchService.UpdateConfig(ctx, &req) if err != nil { common.ErrorWithCode(c, code, err.Error()) return @@ -115,8 +113,6 @@ func (h *SkillSearchHandler) UpdateConfig(c *gin.Context) { // @Summary Search Skills // @Description Search skills using configured search strategy // @Tags skill-search -// @Accept json -// @Produce json // @Security ApiKeyAuth // @Param request body service.SearchRequest true "search query" // @Success 200 {object} map[string]interface{} @@ -156,8 +152,6 @@ type IndexSkillsRequest struct { // @Summary Index Skills // @Description Index skills for search. If embd_id is not provided, will use the one from skill search config. // @Tags skill-search -// @Accept json -// @Produce json // @Security ApiKeyAuth // @Param request body IndexSkillsRequest true "skills to index" // @Success 200 {object} map[string]interface{} @@ -178,7 +172,8 @@ func (h *SkillSearchHandler) IndexSkills(c *gin.Context) { // If embd_id not provided, get from skill search config embdID := req.EmbdID if embdID == "" { - config, code, err := h.searchService.GetConfig(user.ID, req.SpaceID, "") + ctx := c.Request.Context() + config, code, err := h.searchService.GetConfig(ctx, user.ID, req.SpaceID, "") if err != nil { common.ResponseWithCodeData(c, code, nil, "failed to get skill search config: "+err.Error()) return @@ -231,8 +226,6 @@ type ReindexRequest struct { // @Summary Reindex All Skills // @Description Reindex all skills for a tenant. If embd_id is not provided, will use the one from skill search config. // @Tags skill-search -// @Accept json -// @Produce json // @Security ApiKeyAuth // @Param request body ReindexRequest true "skills to reindex" // @Success 200 {object} map[string]interface{} @@ -253,7 +246,8 @@ func (h *SkillSearchHandler) Reindex(c *gin.Context) { // If embd_id not provided, get from skill search config embdID := req.EmbdID if embdID == "" { - config, code, err := h.searchService.GetConfig(user.ID, req.SpaceID, "") + ctx := c.Request.Context() + config, code, err := h.searchService.GetConfig(ctx, user.ID, req.SpaceID, "") if err != nil { common.ResponseWithCodeData(c, code, nil, "failed to get skill search config: "+err.Error()) return @@ -348,8 +342,6 @@ func (h *SkillSearchHandler) InitializeIndex(c *gin.Context) { // @Summary List Skill Spaces // @Description List all skill spaces for the current tenant // @Tags skill-space -// @Accept json -// @Produce json // @Security ApiKeyAuth // @Success 200 {object} map[string]interface{} // @Router /api/v1/skills/spaces [get] @@ -360,7 +352,8 @@ func (h *SkillSearchHandler) ListSpaces(c *gin.Context) { return } - result, code, err := h.spaceService.ListSpaces(user.ID) + ctx := c.Request.Context() + result, code, err := h.spaceService.ListSpaces(ctx, user.ID) if err != nil { common.ErrorWithCode(c, code, err.Error()) return @@ -420,8 +413,6 @@ func (h *SkillSearchHandler) CreateSpace(c *gin.Context) { // @Summary Get Skill Space // @Description Get a skill space by ID // @Tags skill-space -// @Accept json -// @Produce json // @Security ApiKeyAuth // @Param space_id path string true "Space ID" // @Success 200 {object} map[string]interface{} @@ -439,7 +430,8 @@ func (h *SkillSearchHandler) GetSpace(c *gin.Context) { return } - result, code, err := h.spaceService.GetSpace(spaceID, user.ID) + ctx := c.Request.Context() + result, code, err := h.spaceService.GetSpace(ctx, spaceID, user.ID) if err != nil { common.ErrorWithCode(c, code, err.Error()) return @@ -526,7 +518,7 @@ func (h *SkillSearchHandler) DeleteSpace(c *gin.Context) { return } - code, err := h.spaceService.DeleteSpace(spaceID, user.ID, h.docEngine, c.Request.Context()) + code, err := h.spaceService.DeleteSpace(c.Request.Context(), spaceID, user.ID, h.docEngine) if err != nil { common.ErrorWithCode(c, code, err.Error()) return @@ -563,7 +555,8 @@ func (h *SkillSearchHandler) GetSpaceByFolder(c *gin.Context) { return } - result, code, err := h.spaceService.GetSpaceByFolderID(folderID, user.ID) + ctx := c.Request.Context() + result, code, err := h.spaceService.GetSpaceByFolderID(ctx, folderID, user.ID) if err != nil { common.ErrorWithCode(c, code, err.Error()) return diff --git a/internal/ingestion/task/pipeline_executor.go b/internal/ingestion/task/pipeline_executor.go index 1f4f1595b8..fc509d0bc3 100644 --- a/internal/ingestion/task/pipeline_executor.go +++ b/internal/ingestion/task/pipeline_executor.go @@ -29,6 +29,8 @@ import ( "ragflow/internal/engine" "ragflow/internal/entity" pipelinepkg "ragflow/internal/ingestion/pipeline" + + "gorm.io/gorm" ) // PipelineResult is the outcome of a pipeline run: chunks have been @@ -49,7 +51,7 @@ type PipelineExecutor struct { docBulkSize int indexWriter *chunkIndexWriter - logCreateFunc func(log *entity.PipelineOperationLog) error + logCreateFunc func(ctx context.Context, db *gorm.DB, log *entity.PipelineOperationLog) error loadDSLFunc func(ctx context.Context, canvasID string) (string, string, error) runPipelineFunc func(ctx context.Context, dsl string) (map[string]any, string, error) progressSink pipelinepkg.ProgressSink @@ -113,7 +115,7 @@ func (s *PipelineExecutor) WithInsertFunc(f InsertFunc) *PipelineExecutor { return s } -func (s *PipelineExecutor) WithLogCreateFunc(f func(log *entity.PipelineOperationLog) error) *PipelineExecutor { +func (s *PipelineExecutor) WithLogCreateFunc(f func(ctx context.Context, db *gorm.DB, log *entity.PipelineOperationLog) error) *PipelineExecutor { s.logCreateFunc = f return s } @@ -168,7 +170,7 @@ func (s *PipelineExecutor) Execute(ctx context.Context) (*PipelineResult, error) } if s.taskCtx.Doc.ID == CANVAS_DEBUG_DOC_ID { - s.recordPipelineLog(s.taskCtx.Doc.ID, pipelineDSL, "done") + s.recordPipelineLog(ctx, dao.DB, s.taskCtx.Doc.ID, pipelineDSL, "done") return nil, nil } @@ -178,7 +180,7 @@ func (s *PipelineExecutor) Execute(ctx context.Context) (*PipelineResult, error) } if pipelineDSL != "" { - s.recordPipelineLog(s.taskCtx.Doc.ID, pipelineDSL, "done") + s.recordPipelineLog(ctx, dao.DB, s.taskCtx.Doc.ID, pipelineDSL, "done") } return result, nil } @@ -254,7 +256,7 @@ func countDistinctChunkIDs(chunks []map[string]any) int { return len(seen) } -func (s *PipelineExecutor) recordPipelineLog(docID, dsl, status string) { +func (s *PipelineExecutor) recordPipelineLog(ctx context.Context, db *gorm.DB, docID, dsl, status string) { var dslMap entity.JSONMap if err := json.Unmarshal([]byte(dsl), &dslMap); err != nil { dslMap = entity.JSONMap{"raw": dsl} @@ -274,7 +276,7 @@ func (s *PipelineExecutor) recordPipelineLog(docID, dsl, status string) { SourceFrom: s.taskCtx.Doc.SourceType, OperationStatus: status, } - if err := s.logCreateFunc(log); err != nil { + if err := s.logCreateFunc(ctx, db, log); err != nil { common.Warn(fmt.Sprintf("failed to record pipeline log: %v", err)) } } diff --git a/internal/ingestion/task/pipeline_executor_test.go b/internal/ingestion/task/pipeline_executor_test.go index a32a173c01..67692f9ad7 100644 --- a/internal/ingestion/task/pipeline_executor_test.go +++ b/internal/ingestion/task/pipeline_executor_test.go @@ -185,7 +185,8 @@ func TestInsertChunks_EmptyChunks(t *testing.T) { return nil, nil }, ) - err := svc.indexWriter.Write(context.Background(), nil) + ctx := t.Context() + err := svc.indexWriter.Write(ctx, nil) if err != nil { t.Errorf("expected no error for nil chunks, got %v", err) } @@ -200,8 +201,9 @@ func TestInsertChunks_BaseNameAndDatasetID(t *testing.T) { return nil, nil }, ) + ctx := t.Context() chunks := []map[string]any{{"text": "hello"}} - err := svc.indexWriter.Write(context.Background(), chunks) + err := svc.indexWriter.Write(ctx, chunks) if err != nil { t.Fatalf("unexpected error: %v", err) } @@ -215,20 +217,22 @@ func TestInsertChunks_BaseNameAndDatasetID(t *testing.T) { func TestRecordPipelineLog(t *testing.T) { svc := mustNewPipelineExecutor(t, makeTaskCtx(), "flow-1", 0).WithLogCreateFunc( - func(log *entity.PipelineOperationLog) error { return nil }, + func(ctx context.Context, db *gorm.DB, log *entity.PipelineOperationLog) error { return nil }, ) - svc.recordPipelineLog("doc-1", `{"components": {}}`, "done") + ctx := t.Context() + svc.recordPipelineLog(ctx, dao.DB, "doc-1", `{"components": {}}`, "done") } func TestRecordPipelineLog_InvalidJSONFallback(t *testing.T) { var captured *entity.PipelineOperationLog svc := mustNewPipelineExecutor(t, makeTaskCtx(), "flow-1", 0).WithLogCreateFunc( - func(log *entity.PipelineOperationLog) error { + func(ctx context.Context, db *gorm.DB, log *entity.PipelineOperationLog) error { captured = log return nil }, ) - svc.recordPipelineLog("doc-1", "not-valid-json", "done") + ctx := t.Context() + svc.recordPipelineLog(ctx, dao.DB, "doc-1", "not-valid-json", "done") if captured == nil { t.Fatal("logCreateFunc was not called") } @@ -241,12 +245,13 @@ func TestRecordPipelineLog_InvalidJSONFallback(t *testing.T) { func TestRecordPipelineLog_ValidJSONParsed(t *testing.T) { var captured *entity.PipelineOperationLog svc := mustNewPipelineExecutor(t, makeTaskCtx(), "flow-1", 0).WithLogCreateFunc( - func(log *entity.PipelineOperationLog) error { + func(ctx context.Context, db *gorm.DB, log *entity.PipelineOperationLog) error { captured = log return nil }, ) - svc.recordPipelineLog("doc-1", `{"components": {"a": {"obj": {"component_name": "Parser", "params": {}}}}}`, "done") + ctx := t.Context() + svc.recordPipelineLog(ctx, dao.DB, "doc-1", `{"components": {"a": {"obj": {"component_name": "Parser", "params": {}}}}}`, "done") if captured == nil { t.Fatal("logCreateFunc was not called") } @@ -261,7 +266,8 @@ func TestRecordPipelineLog_ValidJSONParsed(t *testing.T) { func TestRunPipeline_NilOutput(t *testing.T) { svc := mustNewPipelineExecutor(t, makeTaskCtx(), "flow-1", 0) - _, err := svc.processOutput(context.Background(), nil, time.Now()) + ctx := t.Context() + _, err := svc.processOutput(ctx, nil, time.Now()) if err != nil { t.Errorf("expected nil error for nil output, got %v", err) } @@ -269,9 +275,10 @@ func TestRunPipeline_NilOutput(t *testing.T) { func TestRunPipeline_EmptyOutput(t *testing.T) { svc := mustNewPipelineExecutor(t, makeTaskCtx(), "flow-1", 0).WithLogCreateFunc( - func(log *entity.PipelineOperationLog) error { return nil }, + func(ctx context.Context, db *gorm.DB, log *entity.PipelineOperationLog) error { return nil }, ) - _, err := svc.processOutput(context.Background(), map[string]any{}, time.Now()) + ctx := t.Context() + _, err := svc.processOutput(ctx, map[string]any{}, time.Now()) if err != nil { t.Errorf("expected nil error for empty output, got %v", err) } @@ -279,9 +286,10 @@ func TestRunPipeline_EmptyOutput(t *testing.T) { func TestRunPipeline_NormalizedEmpty(t *testing.T) { svc := mustNewPipelineExecutor(t, makeTaskCtx(), "flow-1", 0).WithLogCreateFunc( - func(log *entity.PipelineOperationLog) error { return nil }, + func(ctx context.Context, db *gorm.DB, log *entity.PipelineOperationLog) error { return nil }, ) - _, err := svc.processOutput(context.Background(), map[string]any{"markdown": ""}, time.Now()) + ctx := t.Context() + _, err := svc.processOutput(ctx, map[string]any{"markdown": ""}, time.Now()) if err != nil { t.Errorf("expected nil error for empty normalized output, got %v", err) } @@ -292,14 +300,15 @@ func TestRunPipeline_FullFlow(t *testing.T) { WithInsertFunc(func(ctx context.Context, chunks []map[string]any, baseName, datasetID string) ([]string, error) { return nil, nil }). - WithLogCreateFunc(func(log *entity.PipelineOperationLog) error { return nil }) + WithLogCreateFunc(func(ctx context.Context, db *gorm.DB, log *entity.PipelineOperationLog) error { return nil }) output := map[string]any{ "chunks": []map[string]any{ {"text": "hello"}, {"text": "world"}, }, } - _, err := svc.processOutput(context.Background(), output, time.Now()) + ctx := t.Context() + _, err := svc.processOutput(ctx, output, time.Now()) if err != nil { t.Errorf("unexpected error: %v", err) } @@ -310,14 +319,15 @@ func TestRunPipeline_AlreadyHasVectors(t *testing.T) { WithInsertFunc(func(ctx context.Context, chunks []map[string]any, baseName, datasetID string) ([]string, error) { return nil, nil }). - WithLogCreateFunc(func(log *entity.PipelineOperationLog) error { return nil }) + WithLogCreateFunc(func(ctx context.Context, db *gorm.DB, log *entity.PipelineOperationLog) error { return nil }) output := map[string]any{ "chunks": []map[string]any{ {"text": "hello", "q_768_vec": []float64{0.1, 0.2}}, }, } - _, err := svc.processOutput(context.Background(), output, time.Now()) + ctx := t.Context() + _, err := svc.processOutput(ctx, output, time.Now()) if err != nil { t.Errorf("unexpected error: %v", err) } @@ -355,7 +365,7 @@ func TestPipelineExecutor_Run_MainFlowWithStubs(t *testing.T) { inserted = true return nil, nil }). - WithLogCreateFunc(func(log *entity.PipelineOperationLog) error { + WithLogCreateFunc(func(ctx context.Context, db *gorm.DB, log *entity.PipelineOperationLog) error { logged = true if log.PipelineID == nil || *log.PipelineID != "flow-corrected" { t.Fatalf("PipelineID = %v, want flow-corrected", log.PipelineID) @@ -382,7 +392,8 @@ func TestPipelineExecutor_Execute_PropagatesContext(t *testing.T) { type ctxKey string const key ctxKey = "trace" taskCtx := makeTaskCtx() - taskCtx.Ctx = context.WithValue(context.Background(), key, "task-ctx") + ctx := t.Context() + taskCtx.Ctx = context.WithValue(ctx, key, "task-ctx") svc := mustNewPipelineExecutor(t, taskCtx, "flow-1", 0). WithLoadDSLFunc(func(ctx context.Context, canvasID string) (string, string, error) { @@ -397,7 +408,7 @@ func TestPipelineExecutor_Execute_PropagatesContext(t *testing.T) { WithInsertFunc(func(ctx context.Context, chunks []map[string]any, baseName, datasetID string) ([]string, error) { return nil, nil }). - WithLogCreateFunc(func(log *entity.PipelineOperationLog) error { return nil }) + WithLogCreateFunc(func(ctx context.Context, db *gorm.DB, log *entity.PipelineOperationLog) error { return nil }) if _, err := svc.Execute(taskCtx.Ctx); err != nil { t.Fatalf("unexpected error: %v", err) @@ -450,8 +461,9 @@ func TestPipelineExecutorRunPipelineWithDSLForwardsSink(t *testing.T) { svc.WithProgressSink(sink) dsl := `{"dsl":{"components":{"begin":{"obj":{"component_name":"Begin","params":{}},"downstream":["a"]},"a":{"obj":{"component_name":"` + nameA + `","params":{}},"upstream":["begin"]}},"path":["begin","a"],"graph":{"nodes":[]}}}` + ctx := t.Context() - if _, _, err := svc.runPipelineWithDSL(context.Background(), dsl); err != nil { + if _, _, err := svc.runPipelineWithDSL(ctx, dsl); err != nil { t.Fatalf("runPipelineWithDSL: %v", err) } diff --git a/internal/ingestion/task/pipeline_real_integration_test.go b/internal/ingestion/task/pipeline_real_integration_test.go index 5917b9a229..d047271873 100644 --- a/internal/ingestion/task/pipeline_real_integration_test.go +++ b/internal/ingestion/task/pipeline_real_integration_test.go @@ -162,7 +162,7 @@ func TestPipelineExecutor_Run_RealPDF_ProducesIndexedChunks(t *testing.T) { } templateBytes = disableTokenizerEmbeddingForTaskTemplate(t, templateBytes) var templateDSL entity.JSONMap - if err := json.Unmarshal(templateBytes, &templateDSL); err != nil { + if err = json.Unmarshal(templateBytes, &templateDSL); err != nil { t.Fatalf("unmarshal template dsl: %v", err) } @@ -184,10 +184,10 @@ func TestPipelineExecutor_Run_RealPDF_ProducesIndexedChunks(t *testing.T) { objectPath := fmt.Sprintf("integration/task/%s/%s", docID, docName) mustSeedTaskRealPipelineDocumentBytes(t, realDB, realStorage, tenantID, kbID, docID, fileID, bucket, objectPath, docName, ".pdf", "pdf", pdfBytes) - if err := realDB.Model(&entity.Document{}).Where("id = ?", docID).Update("pipeline_id", canvasID).Error; err != nil { + if err = realDB.Model(&entity.Document{}).Where("id = ?", docID).Update("pipeline_id", canvasID).Error; err != nil { t.Fatalf("set document pipeline_id: %v", err) } - if err := realDB.Create(&entity.UserCanvas{ + if err = realDB.Create(&entity.UserCanvas{ ID: canvasID, UserID: tenantID, Permission: "me", @@ -227,9 +227,10 @@ func TestPipelineExecutor_Run_RealPDF_ProducesIndexedChunks(t *testing.T) { inserted = append(inserted, deepCopyTaskChunks(chunks)) return nil, nil }). - WithLogCreateFunc(func(log *entity.PipelineOperationLog) error { return nil }) + WithLogCreateFunc(func(ctx context.Context, db *gorm.DB, log *entity.PipelineOperationLog) error { return nil }) - if _, err := svc.Execute(context.Background()); err != nil { + ctx := t.Context() + if _, err = svc.Execute(ctx); err != nil { t.Fatalf("Run: %v", err) } @@ -311,7 +312,8 @@ func TestRunPipeline_RealPipelineOutput_ProducesIndexFields(t *testing.T) { cleanupTaskRealPipelineDocument(realDB, realStorage, tenantID, kbID, docID, fileID, bucket, objectPath) }) - pipelineOut, err := pipe.Run(context.Background(), map[string]any{ + ctx := t.Context() + pipelineOut, err := pipe.Run(ctx, map[string]any{ "doc_id": docID, }, nil) if err != nil { @@ -345,7 +347,7 @@ func TestRunPipeline_RealPipelineOutput_ProducesIndexFields(t *testing.T) { return nil, nil }) - if _, err := svc.processOutput(context.Background(), pipelineOut, time.Now()); err != nil { + if _, err = svc.processOutput(ctx, pipelineOut, time.Now()); err != nil { t.Fatalf("RunPipeline: %v", err) } diff --git a/internal/service/chat.go b/internal/service/chat.go index 869a78f8de..be5b7b51c3 100644 --- a/internal/service/chat.go +++ b/internal/service/chat.go @@ -235,7 +235,7 @@ func (s *ChatService) Create(ctx context.Context, userID string, req map[string] if llmIDValue, ok := req["llm_id"]; ok { llmID := stringFromValue(llmIDValue) llmSetting, _ := mapFromValue(req["llm_setting"]) - tenantLLMID, err := resolveCreateLLMID(llmID, userID, llmSetting) + tenantLLMID, err := resolveCreateLLMID(ctx, llmID, userID, llmSetting) if err != nil { return nil, common.CodeDataError, err } @@ -246,7 +246,7 @@ func (s *ChatService) Create(ctx context.Context, userID string, req map[string] if rerankIDValue, ok := req["rerank_id"]; ok { rerankID := stringFromValue(rerankIDValue) - tenantRerankID, err := resolveCreateRerankID(rerankID, userID) + tenantRerankID, err := resolveCreateRerankID(ctx, rerankID, userID) if err != nil { return nil, common.CodeDataError, err } @@ -278,7 +278,7 @@ func (s *ChatService) Create(ctx context.Context, userID string, req map[string] } if stringFromValue(req["llm_id"]) != "" && !isTruthy(req["tenant_llm_id"]) { llmSetting, _ := mapFromValue(req["llm_setting"]) - tenantLLMID, err := resolveCreateLLMID(stringFromValue(req["llm_id"]), userID, llmSetting) + tenantLLMID, err := resolveCreateLLMID(ctx, stringFromValue(req["llm_id"]), userID, llmSetting) if err != nil { return nil, common.CodeDataError, err } @@ -403,7 +403,7 @@ func (s *ChatService) validateCreateDatasetIDs(ctx context.Context, value interf return normalizedIDs, nil } -func resolveCreateLLMID(llmID, tenantID string, llmSetting map[string]interface{}) (string, error) { +func resolveCreateLLMID(ctx context.Context, llmID, tenantID string, llmSetting map[string]interface{}) (string, error) { if llmID == "" { return "", nil } @@ -429,17 +429,17 @@ func resolveCreateLLMID(llmID, tenantID string, llmSetting map[string]interface{ } } modelProvider := NewModelProviderService() - if _, _, _, _, err := modelProvider.ResolveModelConfig(tenantID, modelType, llmID); err != nil { + if _, _, _, _, err := modelProvider.ResolveModelConfig(ctx, tenantID, modelType, llmID); err != nil { return "", fmt.Errorf("`llm_id` %s doesn't exist", llmID) } - tenantLLMID, err := modelProvider.ResolveModelID(tenantID, modelType, llmID) + tenantLLMID, err := modelProvider.ResolveModelID(ctx, tenantID, modelType, llmID) if err != nil { return "", err } return tenantLLMID, nil } -func resolveCreateRerankID(rerankID, tenantID string) (string, error) { +func resolveCreateRerankID(ctx context.Context, rerankID, tenantID string) (string, error) { if rerankID == "" { return "", nil } @@ -448,10 +448,10 @@ func resolveCreateRerankID(rerankID, tenantID string) (string, error) { return "", nil } modelProvider := NewModelProviderService() - if _, _, _, _, err := modelProvider.ResolveModelConfig(tenantID, entity.ModelTypeRerank, rerankID); err != nil { + if _, _, _, _, err := modelProvider.ResolveModelConfig(ctx, tenantID, entity.ModelTypeRerank, rerankID); err != nil { return "", fmt.Errorf("`rerank_id` %s doesn't exist", rerankID) } - tenantRerankID, err := modelProvider.ResolveModelID(tenantID, entity.ModelTypeRerank, rerankID) + tenantRerankID, err := modelProvider.ResolveModelID(ctx, tenantID, entity.ModelTypeRerank, rerankID) if err != nil { return "", err } @@ -890,7 +890,7 @@ func (s *ChatService) updateChatREST(ctx context.Context, userID, chatID string, if value, ok := req["llm_id"]; ok { llmID := fmt.Sprint(value) - tenantLLMID, err := s.resolveRESTLLMID(llmID, userID, llmSetting) + tenantLLMID, err := s.resolveRESTLLMID(ctx, llmID, userID, llmSetting) if err != nil { return nil, err } @@ -901,7 +901,7 @@ func (s *ChatService) updateChatREST(ctx context.Context, userID, chatID string, if value, ok := req["rerank_id"]; ok { rerankID := fmt.Sprint(value) - tenantRerankID, err := s.resolveRESTRerankID(rerankID, userID) + tenantRerankID, err := s.resolveRESTRerankID(ctx, rerankID, userID) if err != nil { return nil, err } @@ -1046,7 +1046,7 @@ func (s *ChatService) validateRESTDatasetIDs(ctx context.Context, value interfac return kbIDs, nil } -func (s *ChatService) resolveRESTLLMID(llmID, tenantID string, llmSetting map[string]interface{}) (string, error) { +func (s *ChatService) resolveRESTLLMID(ctx context.Context, llmID, tenantID string, llmSetting map[string]interface{}) (string, error) { if llmID == "" { return "", nil } @@ -1067,17 +1067,17 @@ func (s *ChatService) resolveRESTLLMID(llmID, tenantID string, llmSetting map[st } } modelProvider := NewModelProviderService() - if _, _, _, _, err := modelProvider.ResolveModelConfig(tenantID, modelType, llmID); err != nil { + if _, _, _, _, err := modelProvider.ResolveModelConfig(ctx, tenantID, modelType, llmID); err != nil { return "", fmt.Errorf("`llm_id` %s doesn't exist", llmID) } - tenantLLMID, err := modelProvider.ResolveModelID(tenantID, modelType, llmID) + tenantLLMID, err := modelProvider.ResolveModelID(ctx, tenantID, modelType, llmID) if err != nil { return "", err } return tenantLLMID, nil } -func (s *ChatService) resolveRESTRerankID(rerankID, tenantID string) (string, error) { +func (s *ChatService) resolveRESTRerankID(ctx context.Context, rerankID, tenantID string) (string, error) { if rerankID == "" { return "", nil } @@ -1086,10 +1086,10 @@ func (s *ChatService) resolveRESTRerankID(rerankID, tenantID string) (string, er return "", nil } modelProvider := NewModelProviderService() - if _, _, _, _, err := modelProvider.ResolveModelConfig(tenantID, entity.ModelTypeRerank, rerankID); err != nil { + if _, _, _, _, err := modelProvider.ResolveModelConfig(ctx, tenantID, entity.ModelTypeRerank, rerankID); err != nil { return "", fmt.Errorf("`rerank_id` %s doesn't exist", rerankID) } - tenantRerankID, err := modelProvider.ResolveModelID(tenantID, entity.ModelTypeRerank, rerankID) + tenantRerankID, err := modelProvider.ResolveModelID(ctx, tenantID, entity.ModelTypeRerank, rerankID) if err != nil { return "", err } diff --git a/internal/service/chat_pipeline.go b/internal/service/chat_pipeline.go index ce5ed22b9c..5c311eebbe 100644 --- a/internal/service/chat_pipeline.go +++ b/internal/service/chat_pipeline.go @@ -169,7 +169,7 @@ func (s *ChatPipelineService) AsyncChat( } lastMsg := messages[len(messages)-1] if role, _ := lastMsg["role"].(string); role != "user" { - return nil, fmt.Errorf("The last content of this conversation is not from user.") + return nil, fmt.Errorf("the last content of this conversation is not from user") } // No KBs & no web search → fast-path to LLM-only chat. @@ -205,7 +205,7 @@ func (s *ChatPipelineService) AsyncChat( // === Phase 2: Resolve LLM Model Config + max_tokens === common.Info("Phase 2: Resolve LLM Model Config + max_tokens") timer.Enter(common.PhaseCheckLLM) - llmModelConfig, _, _, _, err := s.getLLMModelConfig(chat) + llmModelConfig, _, _, _, err := s.getLLMModelConfig(ctx, chat) if err != nil { out <- AsyncChatResult{ Answer: fmt.Sprintf("**ERROR**: %s", err.Error()), @@ -1009,7 +1009,7 @@ func (s *ChatPipelineService) AsyncChat( zap.Bool("stream", stream), zap.Int("llm_messages_count", len(llmMessages))) timer.Enter(common.PhaseGenerateAnswer) - chatDriver := s.buildChatDriver(chat, chatModel) + chatDriver := s.buildChatDriver(ctx, chat, chatModel) if chatDriver == nil { out <- AsyncChatResult{ Answer: "**ERROR**: No chat model available for this chat.", @@ -1329,7 +1329,7 @@ func (s *ChatPipelineService) AsyncChatSolo( } // 1b. Resolve LLM model config (needed early for model_type dispatch). - llmModelConfig, _, _, _, err := s.getLLMModelConfig(chat) + llmModelConfig, _, _, _, err := s.getLLMModelConfig(ctx, chat) factoryName := "" if err == nil && llmModelConfig != nil { factoryName, _ = llmModelConfig["llm_factory"].(string) @@ -1383,7 +1383,7 @@ func (s *ChatPipelineService) AsyncChatSolo( } // 4. Build the chat model wrapper. - driver, modelName, apiConfig, _, err := s.ModelProviderSvc.GetChatModelConfig(chat.TenantID, chat.LLMID) + driver, modelName, apiConfig, _, err := s.ModelProviderSvc.GetChatModelConfig(ctx, chat.TenantID, chat.LLMID) if err != nil { out <- AsyncChatResult{ Answer: fmt.Sprintf("**ERROR**: %s", err.Error()), @@ -1398,7 +1398,7 @@ func (s *ChatPipelineService) AsyncChatSolo( if promptConfig != nil { if useTTS, _ := promptConfig["tts"].(bool); useTTS { ttsDriver, ttsName, ttsConfig, _, ttsErr := s.ModelProviderSvc.GetTenantDefaultModelByType( - chat.TenantID, entity.ModelTypeTTS, + ctx, chat.TenantID, entity.ModelTypeTTS, ) if ttsErr != nil { common.Warn("AsyncChatSolo: TTS lookup failed; proceeding without TTS", @@ -1876,11 +1876,11 @@ func tokenizeText(text string) string { // The returned `cfg` map's "model_type" field carries the chosen type // so downstream code (e.g. the multimodal-conversion guard in AsyncChat // at async_chat.go:632) can skip chat-only logic for image2text dialogs. -func (s *ChatPipelineService) getLLMModelConfig(chat *entity.Chat) (map[string]interface{}, string, string, string, error) { +func (s *ChatPipelineService) getLLMModelConfig(ctx context.Context, chat *entity.Chat) (map[string]interface{}, string, string, string, error) { if chat.LLMID == "" { // Branch 3: no explicit LLM → tenant default chat model. return s.buildLLMModelConfig( - s.ModelProviderSvc.GetTenantDefaultModelByType(chat.TenantID, entity.ModelTypeChat), + s.ModelProviderSvc.GetTenantDefaultModelByType(ctx, chat.TenantID, entity.ModelTypeChat), ) } @@ -1898,7 +1898,7 @@ func (s *ChatPipelineService) getLLMModelConfig(chat *entity.Chat) (map[string]i } } cfg, modelName, factoryName, baseURL, err := s.buildLLMModelConfig( - s.ModelProviderSvc.ResolveModelConfig(chat.TenantID, modelType, chat.LLMID), + s.ModelProviderSvc.ResolveModelConfig(ctx, chat.TenantID, modelType, chat.LLMID), ) if err != nil { return nil, "", "", "", err @@ -1978,7 +1978,7 @@ func (s *ChatPipelineService) getModels(ctx context.Context, chat *entity.Chat) if kbs[0].EmbdID != "" { embdTenantID := kbs[0].TenantID driver, modelName, apiConfig, maxTokens, err := s.ModelProviderSvc.ResolveModelConfig( - embdTenantID, entity.ModelTypeEmbedding, kbs[0].EmbdID, + ctx, embdTenantID, entity.ModelTypeEmbedding, kbs[0].EmbdID, ) if err != nil { common.Warn("Failed to get embedding model for chat retrieval", @@ -1992,7 +1992,7 @@ func (s *ChatPipelineService) getModels(ctx context.Context, chat *entity.Chat) } // Chat model. - driver, modelName, apiConfig, _, err := s.ModelProviderSvc.GetChatModelConfig(chat.TenantID, chat.LLMID) + driver, modelName, apiConfig, _, err := s.ModelProviderSvc.GetChatModelConfig(ctx, chat.TenantID, chat.LLMID) var chatModel *modelModule.ChatModel if err == nil { chatModel = modelModule.NewChatModel(driver, &modelName, apiConfig) @@ -2002,7 +2002,7 @@ func (s *ChatPipelineService) getModels(ctx context.Context, chat *entity.Chat) var rerankModel *modelModule.RerankModel if chat.RerankID != "" { rerankDriver, rerankName, rerankConfig, _, err := s.ModelProviderSvc.ResolveModelConfig( - chat.TenantID, entity.ModelTypeRerank, chat.RerankID, + ctx, chat.TenantID, entity.ModelTypeRerank, chat.RerankID, ) if err == nil { rerankModel = modelModule.NewRerankModel(rerankDriver, &rerankName, rerankConfig) @@ -2014,7 +2014,7 @@ func (s *ChatPipelineService) getModels(ctx context.Context, chat *entity.Chat) if chat.PromptConfig != nil { if useTTS, _ := chat.PromptConfig["tts"].(bool); useTTS { ttsDriver, ttsName, ttsConfig, _, err := s.ModelProviderSvc.GetTenantDefaultModelByType( - chat.TenantID, entity.ModelTypeTTS, + ctx, chat.TenantID, entity.ModelTypeTTS, ) if err == nil { ttsModel = modelModule.NewChatModel(ttsDriver, &ttsName, ttsConfig) @@ -2554,11 +2554,11 @@ func (s *ChatPipelineService) buildChatMessages(systemContent string, messages [ } // buildChatDriver creates a ChatModel wrapper from the chat. -func (s *ChatPipelineService) buildChatDriver(chat *entity.Chat, chatModel *modelModule.ChatModel) *modelModule.ChatModel { +func (s *ChatPipelineService) buildChatDriver(ctx context.Context, chat *entity.Chat, chatModel *modelModule.ChatModel) *modelModule.ChatModel { if chatModel != nil { return chatModel } - driver, modelName, apiConfig, _, err := s.ModelProviderSvc.GetChatModelConfig(chat.TenantID, chat.LLMID) + driver, modelName, apiConfig, _, err := s.ModelProviderSvc.GetChatModelConfig(ctx, chat.TenantID, chat.LLMID) if err != nil { return nil } @@ -3710,7 +3710,7 @@ func removeRedundantSpaces(s string) string { // dialog_service.py:1309. Matches `T13:24:55|` or `T13:24:55.123Z|`. var isoTimestampCellRe = regexp.MustCompile(`T[0-9]{2}:[0-9]{2}:[0-9]{2}(\.[0-9]+Z)?\|`) -// stripISOTimestamps removes ISO-8601 timestamps that end a markdown +// stripISOTimestamps removes ISO-8601 timestamps that end a Markdown // table cell. Operates on the full joined rows string (not per-cell). func stripISOTimestamps(rows string) string { return isoTimestampCellRe.ReplaceAllString(rows, "|") diff --git a/internal/service/chat_session.go b/internal/service/chat_session.go index 244cee331a..490c14ce49 100644 --- a/internal/service/chat_session.go +++ b/internal/service/chat_session.go @@ -60,7 +60,7 @@ type chatPipelineRunner interface { } type chatModelConfigResolver interface { - GetChatModelConfig(tenantID, llmID string) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) + GetChatModelConfig(ctx context.Context, tenantID, llmID string) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) } // chunkFeedbackApplier is the dispatch seam for chunk-level feedback @@ -1289,7 +1289,7 @@ func (s *ChatSessionService) Completion(ctx context.Context, userID string, conv isEmbedded := llmID != "" if llmID != "" { - hasKey, err := s.checkTenantLLMAPIKey(dialog.TenantID, llmID) + hasKey, err := s.checkTenantLLMAPIKey(ctx, dialog.TenantID, llmID) if err != nil || !hasKey { return nil, fmt.Errorf("Cannot use specified model %s", llmID) } @@ -1378,7 +1378,7 @@ func (s *ChatSessionService) CompletionStream(ctx context.Context, userID string isEmbedded := llmID != "" if llmID != "" { - hasKey, err := s.checkTenantLLMAPIKey(dialog.TenantID, llmID) + hasKey, err := s.checkTenantLLMAPIKey(ctx, dialog.TenantID, llmID) if err != nil || !hasKey { errMsg := fmt.Sprintf(`{"code": 500, "message": "Cannot use specified model %s", "data": {"answer": "**ERROR**: Cannot use specified model", "reference": []}}`, llmID) streamChan <- fmt.Sprintf("data: %s\n\n", errMsg) @@ -1539,7 +1539,7 @@ func (s *ChatSessionService) ChatCompletions( genConfig = map[string]interface{}{} } if llmID != "" { - hasKey, err := s.checkTenantLLMAPIKey(dialog.TenantID, llmID) + hasKey, err := s.checkTenantLLMAPIKey(ctx, dialog.TenantID, llmID) if err != nil || !hasKey { return fail(fmt.Errorf("cannot use specified model %s", llmID)) } @@ -1952,12 +1952,12 @@ func (s *ChatSessionService) initializeReference(session *entity.ChatSession) [] return filtered } -func (s *ChatSessionService) checkTenantLLMAPIKey(tenantID, modelName string) (bool, error) { +func (s *ChatSessionService) checkTenantLLMAPIKey(ctx context.Context, tenantID, modelName string) (bool, error) { resolver := s.modelProviderSvc if resolver == nil { resolver = NewModelProviderService() } - _, _, _, _, err := resolver.GetChatModelConfig(tenantID, modelName) + _, _, _, _, err := resolver.GetChatModelConfig(ctx, tenantID, modelName) if err != nil { return false, err } diff --git a/internal/service/chat_session_test.go b/internal/service/chat_session_test.go index 95f5b9362c..2161cb3845 100644 --- a/internal/service/chat_session_test.go +++ b/internal/service/chat_session_test.go @@ -183,7 +183,7 @@ type fakeChatModelConfigResolver struct { err error } -func (f *fakeChatModelConfigResolver) GetChatModelConfig(tenantID, llmID string) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { +func (f *fakeChatModelConfigResolver) GetChatModelConfig(ctx context.Context, tenantID, llmID string) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { f.tenantID = tenantID f.llmID = llmID if f.err != nil { diff --git a/internal/service/chunk/chunk.go b/internal/service/chunk/chunk.go index 1820ecec22..5376c31db1 100644 --- a/internal/service/chunk/chunk.go +++ b/internal/service/chunk/chunk.go @@ -209,7 +209,7 @@ func (s *ChunkService) RetrievalTest(ctx context.Context, req *service.Retrieval if req.SearchID != nil && *req.SearchID != "" { // If search_id is set, get meta_data_filter and chat_id from search_config - searchDetail, err := s.searchService.GetDetail(*req.SearchID) + searchDetail, err := s.searchService.GetDetail(ctx, *req.SearchID) if err != nil { common.Warn("Failed to get search detail for search_id, proceeding without it", zap.String("searchID", *req.SearchID), zap.Error(err)) } else if searchConfig, ok := searchConfigMap(searchDetail["search_config"]); ok && searchConfig != nil { @@ -229,7 +229,7 @@ func (s *ChunkService) RetrievalTest(ctx context.Context, req *service.Retrieval modelProviderSvc := service.NewModelProviderService() if chatID != "" { // Use chat_id from search_config (it's actually the model name) - driver, mdlName, apiConfig, _, getErr := modelProviderSvc.ResolveModelConfig(tenantIDs[0], entity.ModelTypeChat, chatID) + driver, mdlName, apiConfig, _, getErr := modelProviderSvc.ResolveModelConfig(ctx, tenantIDs[0], entity.ModelTypeChat, chatID) if getErr != nil { common.Warn("Failed to get chat model from search_config chat_id, using tenant default", zap.String("chatID", chatID), zap.Error(getErr)) } else { @@ -248,7 +248,7 @@ func (s *ChunkService) RetrievalTest(ctx context.Context, req *service.Retrieval if err != nil || modelName == "" { common.Warn("Failed to get tenant default chat model name for meta_data_filter", zap.Error(err)) } else { - driver, mdlName, apiConfig, _, getErr := modelProviderSvc.ResolveModelConfig(tenantIDs[0], entity.ModelTypeChat, modelName) + driver, mdlName, apiConfig, _, getErr := modelProviderSvc.ResolveModelConfig(ctx, tenantIDs[0], entity.ModelTypeChat, modelName) if getErr != nil { common.Warn("Failed to get chat model for meta_data_filter", zap.Error(getErr)) } else { @@ -294,7 +294,7 @@ func (s *ChunkService) RetrievalTest(ctx context.Context, req *service.Retrieval if err != nil || llmModelName == "" { common.Warn("Failed to get default chat model name for LLM transformations", zap.Error(err)) } else { - driver, mdlName, apiConfig, _, getErr := modelProviderSvc.ResolveModelConfig(tenantIDs[0], entity.ModelTypeChat, llmModelName) + driver, mdlName, apiConfig, _, getErr := modelProviderSvc.ResolveModelConfig(ctx, tenantIDs[0], entity.ModelTypeChat, llmModelName) if getErr != nil { common.Warn("Failed to get chat model for LLM transformations", zap.Error(getErr)) } else { @@ -344,27 +344,27 @@ func (s *ChunkService) RetrievalTest(ctx context.Context, req *service.Retrieval var embeddingModel *models.EmbeddingModel var embdID string if kbRecords[0].TenantEmbdID != nil && *kbRecords[0].TenantEmbdID != "" { - driver, modelName, apiConfig, maxTokens, getErr := modelProviderSvc.GetModelConfigByID(tenantIDs[0], entity.ModelTypeEmbedding, *kbRecords[0].TenantEmbdID) + driver, modelName, apiConfig, maxTokens, getErr := modelProviderSvc.GetModelConfigByID(ctx, tenantIDs[0], entity.ModelTypeEmbedding, *kbRecords[0].TenantEmbdID) if getErr != nil { return nil, fmt.Errorf("failed to get embedding model by tenant_embd_id: %w", getErr) } embeddingModel = models.NewEmbeddingModel(driver, &modelName, apiConfig, maxTokens) } else if kbRecords[0].EmbdID != "" { embdID = kbRecords[0].EmbdID - driver, modelName, apiConfig, maxTokens, getErr := modelProviderSvc.ResolveModelConfig(tenantIDs[0], entity.ModelTypeEmbedding, embdID) + driver, modelName, apiConfig, maxTokens, getErr := modelProviderSvc.ResolveModelConfig(ctx, tenantIDs[0], entity.ModelTypeEmbedding, embdID) if getErr != nil { _, embdID, err = dao.LookupTenantLLMByName(dao.NewTenantLLMDAO(), tenantIDs[0], kbRecords[0].EmbdID, entity.ModelTypeEmbedding) if err != nil { return nil, fmt.Errorf("failed to get embedding model by embd_id: %w", getErr) } - driver, modelName, apiConfig, maxTokens, getErr = modelProviderSvc.ResolveModelConfig(tenantIDs[0], entity.ModelTypeEmbedding, embdID) + driver, modelName, apiConfig, maxTokens, getErr = modelProviderSvc.ResolveModelConfig(ctx, tenantIDs[0], entity.ModelTypeEmbedding, embdID) if getErr != nil { return nil, fmt.Errorf("failed to get embedding model by embd_id: %w", getErr) } } embeddingModel = models.NewEmbeddingModel(driver, &modelName, apiConfig, maxTokens) } else { - driver, modelName, apiConfig, maxTokens, getErr := modelProviderSvc.GetTenantDefaultModelByType(tenantIDs[0], entity.ModelTypeEmbedding) + driver, modelName, apiConfig, maxTokens, getErr := modelProviderSvc.GetTenantDefaultModelByType(ctx, tenantIDs[0], entity.ModelTypeEmbedding) if getErr != nil { return nil, fmt.Errorf("failed to get tenant default embedding model: %w", getErr) } @@ -383,14 +383,14 @@ func (s *ChunkService) RetrievalTest(ctx context.Context, req *service.Retrieval // Get rerank model if RerankID is specified var rerankModel *models.RerankModel if req.TenantRerankID != nil && *req.TenantRerankID != "" { - driver, mdlName, apiConfig, _, getErr := modelProviderSvc.GetModelConfigByID(tenantIDs[0], entity.ModelTypeRerank, *req.TenantRerankID) + driver, mdlName, apiConfig, _, getErr := modelProviderSvc.GetModelConfigByID(ctx, tenantIDs[0], entity.ModelTypeRerank, *req.TenantRerankID) if getErr != nil { return nil, fmt.Errorf("failed to get rerank model by tenant_rerank_id: %w", getErr) } rerankModel = models.NewRerankModel(driver, &mdlName, apiConfig) } else if req.RerankID != nil && *req.RerankID != "" { rerankCompositeName := *req.RerankID - driver, mdlName, apiConfig, _, getErr := modelProviderSvc.ResolveModelConfig(tenantIDs[0], entity.ModelTypeRerank, rerankCompositeName) + driver, mdlName, apiConfig, _, getErr := modelProviderSvc.ResolveModelConfig(ctx, tenantIDs[0], entity.ModelTypeRerank, rerankCompositeName) if getErr != nil { rerankModel = nil } else { diff --git a/internal/service/dataset/crud.go b/internal/service/dataset/crud.go index fd7dcd1a78..51716e9daf 100644 --- a/internal/service/dataset/crud.go +++ b/internal/service/dataset/crud.go @@ -99,7 +99,7 @@ func (d *DatasetService) CreateDataset(ctx context.Context, req *service.CreateD embdID := tenant.EmbdID tenantEmbdID := ptrStringValue(tenant.TenantEmbdID) if embeddingModel != "" { - ok, message := d.verifyEmbeddingAvailability(embeddingModel, tenantID) + ok, message := d.verifyEmbeddingAvailability(ctx, embeddingModel, tenantID) if !ok { return nil, common.CodeDataError, errors.New(message) } @@ -107,7 +107,7 @@ func (d *DatasetService) CreateDataset(ctx context.Context, req *service.CreateD tenantEmbdID = "" } if embdID != "" && tenantEmbdID == "" { - resolvedID, err := service.NewModelProviderService().ResolveModelID(tenantID, entity.ModelTypeEmbedding, embdID) + resolvedID, err := service.NewModelProviderService().ResolveModelID(ctx, tenantID, entity.ModelTypeEmbedding, embdID) if err == nil { tenantEmbdID = resolvedID } else { diff --git a/internal/service/dataset/index.go b/internal/service/dataset/index.go index 3f7ac1a9de..ed6200734f 100644 --- a/internal/service/dataset/index.go +++ b/internal/service/dataset/index.go @@ -137,7 +137,7 @@ func createDatasetIndexTaskInTx(tx *gorm.DB, task *entity.Task, queueDocID strin func enqueueDatasetIndexTask(priority int, queueMessage map[string]interface{}) error { redisClient := redisengine.Get() if redisClient == nil || !redisClient.QueueProduct(datasetIndexQueueName(priority), queueMessage) { - return errors.New("Can't access Redis. Please check the Redis' status") + return errors.New("can't access Redis. Please check the Redis' status") } return nil } @@ -358,32 +358,32 @@ func (d *DatasetService) getDocumentsByDatasetForIndex(ctx context.Context, data documents, _, err := d.documentDAO.GetByKBID(ctx, dao.DB, datasetID) if err != nil { common.Warn("Failed to load dataset documents for index", zap.String("dataset_id", datasetID), zap.Error(err)) - return nil, common.CodeDataError, errors.New("Internal server error") + return nil, common.CodeDataError, errors.New("internal server error") } if len(documents) == 0 { - return nil, common.CodeDataError, fmt.Errorf("No documents in Dataset %s", datasetID) + return nil, common.CodeDataError, fmt.Errorf("no documents in Dataset %s", datasetID) } return documents, common.CodeSuccess, nil } func (d *DatasetService) TraceIndex(ctx context.Context, datasetID, userID, indexType string) (*entity.Task, common.ErrorCode, error) { if !checkType(indexType) { - return nil, common.CodeDataError, fmt.Errorf("Invalid index type '%s'. Must be one of %v", indexType, validIndexTypes) + return nil, common.CodeDataError, fmt.Errorf("invalid index type '%s'. Must be one of %v", indexType, validIndexTypes) } if datasetID == "" { - return nil, common.CodeDataError, errors.New(`Lack of "Dataset ID"`) + return nil, common.CodeDataError, errors.New(`lack of "Dataset ID"`) } if !d.kbDAO.Accessible(ctx, dao.DB, datasetID, userID) { - return nil, common.CodeDataError, errors.New("No authorization.") + return nil, common.CodeDataError, errors.New("no authorization") } kb, err := d.kbDAO.GetByID(ctx, dao.DB, datasetID) if err != nil { if dao.IsNotFoundErr(err) { - return nil, common.CodeDataError, errors.New("Invalid Dataset ID") + return nil, common.CodeDataError, errors.New("invalid Dataset ID") } - return nil, common.CodeDataError, errors.New("Internal server error") + return nil, common.CodeDataError, errors.New("internal server error") } taskID := datasetIndexTaskID(kb, indexType) @@ -395,7 +395,7 @@ func (d *DatasetService) TraceIndex(ctx context.Context, datasetID, userID, inde if dao.IsNotFoundErr(err) { return nil, common.CodeSuccess, nil } - return nil, common.CodeServerError, errors.New("Internal server error") + return nil, common.CodeServerError, errors.New("internal server error") } if task == nil { return nil, common.CodeSuccess, nil @@ -421,32 +421,32 @@ type embeddingCheckSample struct { func (d *DatasetService) CheckEmbedding(ctx context.Context, userID, datasetID string, req *service.CheckEmbeddingRequest) (*service.EmbeddingCheckResponse, common.ErrorCode, error) { if datasetID == "" { - return nil, common.CodeDataError, errors.New(`Lack of "Dataset ID"`) + return nil, common.CodeDataError, errors.New(`lack of "Dataset ID"`) } if !d.kbDAO.Accessible(ctx, dao.DB, datasetID, userID) { - return nil, common.CodeDataError, errors.New("No authorization.") + return nil, common.CodeDataError, errors.New("no authorization") } kb, err := d.kbDAO.GetByID(ctx, dao.DB, datasetID) if err != nil { if dao.IsNotFoundErr(err) { - return nil, common.CodeDataError, errors.New("Invalid Dataset ID") + return nil, common.CodeDataError, errors.New("invalid Dataset ID") } - return nil, common.CodeServerError, errors.New("Internal server error") + return nil, common.CodeServerError, errors.New("internal server error") } if req == nil || strings.TrimSpace(req.EmbeddingID) == "" { - return nil, common.CodeDataError, errors.New("`embd_id` is required.") + return nil, common.CodeDataError, errors.New("`embd_id` is required") } embeddingID := strings.TrimSpace(req.EmbeddingID) - if ok, message := d.verifyEmbeddingAvailability(embeddingID, userID); !ok { + if ok, message := d.verifyEmbeddingAvailability(ctx, embeddingID, kb.TenantID); !ok { return nil, common.CodeDataError, errors.New(message) } if d.docEngine == nil { return nil, common.CodeServerError, errors.New("doc engine not initialized") } - driver, modelName, apiConfig, maxTokens, err := service.NewModelProviderService().ResolveModelConfig(kb.TenantID, entity.ModelTypeEmbedding, embeddingID) + driver, modelName, apiConfig, maxTokens, err := service.NewModelProviderService().ResolveModelConfig(ctx, kb.TenantID, entity.ModelTypeEmbedding, embeddingID) if err != nil { return nil, common.CodeDataError, err } @@ -654,8 +654,8 @@ func (d *DatasetService) sampleRandomChunksWithVectors(ctx context.Context, tena return samples, nil } -func (d *DatasetService) verifyEmbeddingAvailability(embdID string, tenantID string) (bool, string) { - _, _, _, _, err := service.NewModelProviderService().ResolveModelConfig(tenantID, entity.ModelTypeEmbedding, embdID) +func (d *DatasetService) verifyEmbeddingAvailability(ctx context.Context, embdID string, tenantID string) (bool, string) { + _, _, _, _, err := service.NewModelProviderService().ResolveModelConfig(ctx, tenantID, entity.ModelTypeEmbedding, embdID) if err != nil { return false, err.Error() } @@ -664,23 +664,23 @@ func (d *DatasetService) verifyEmbeddingAvailability(embdID string, tenantID str func (d *DatasetService) DeleteIndex(ctx context.Context, userID, datasetID, indexType string, wipe bool) (common.ErrorCode, error) { if !checkType(indexType) { - return common.CodeArgumentError, fmt.Errorf("Invalid index type '%s'", indexType) + return common.CodeArgumentError, fmt.Errorf("invalid index type '%s'", indexType) } if datasetID == "" { - return common.CodeDataError, errors.New(`Lack of "Dataset ID"`) + return common.CodeDataError, errors.New(`lack of "Dataset ID"`) } if !d.kbDAO.Accessible(ctx, dao.DB, datasetID, userID) { - return common.CodeDataError, errors.New("No authorization.") + return common.CodeDataError, errors.New("no authorization") } kb, err := d.kbDAO.GetByID(ctx, dao.DB, datasetID) if err != nil { if dao.IsNotFoundErr(err) { - return common.CodeDataError, errors.New("Invalid Dataset ID") + return common.CodeDataError, errors.New("invalid Dataset ID") } - return common.CodeDataError, errors.New("Internal server error") + return common.CodeDataError, errors.New("internal server error") } taskFinishAtField := datasetIndexTaskFinishAtColumn(indexType) @@ -695,13 +695,13 @@ func (d *DatasetService) DeleteIndex(ctx context.Context, userID, datasetID, ind } if err := dao.DB.Unscoped().Where("id = ?", taskID).Delete(&entity.Task{}).Error; err != nil { common.Warn("Failed to delete dataset index task", zap.String("dataset_id", datasetID), zap.String("task_id", taskID), zap.Error(err)) - return common.CodeDataError, errors.New("Internal server error") + return common.CodeDataError, errors.New("internal server error") } } if wipe && indexType == "graph" { if d.docEngine == nil { - return common.CodeServerError, errors.New("Document engine is not initialized") + return common.CodeServerError, errors.New("document engine is not initialized") } indexName := fmt.Sprintf("ragflow_%s", kb.TenantID) _, err = d.docEngine.DeleteChunks(ctx, map[string]interface{}{ @@ -710,13 +710,13 @@ func (d *DatasetService) DeleteIndex(ctx context.Context, userID, datasetID, ind }, indexName, datasetID) if err != nil { common.Warn("Failed to delete GraphRAG artefacts", zap.String("dataset_id", datasetID), zap.Error(err)) - return common.CodeDataError, errors.New("Internal server error") + return common.CodeDataError, errors.New("internal server error") } clearGraphPhaseMarkers(redisengine.Get(), datasetID) common.Info("delete_index: cleared GraphRAG artefacts and phase markers", zap.String("dataset_id", datasetID)) } else if wipe && indexType == "raptor" { if d.docEngine == nil { - return common.CodeServerError, errors.New("Document engine is not initialized") + return common.CodeServerError, errors.New("document engine is not initialized") } indexName := fmt.Sprintf("ragflow_%s", kb.TenantID) _, err = d.docEngine.DeleteChunks(ctx, map[string]interface{}{ @@ -725,7 +725,7 @@ func (d *DatasetService) DeleteIndex(ctx context.Context, userID, datasetID, ind }, indexName, datasetID) if err != nil { common.Warn("Failed to delete RAPTOR artefacts", zap.String("dataset_id", datasetID), zap.Error(err)) - return common.CodeDataError, errors.New("Internal server error") + return common.CodeDataError, errors.New("internal server error") } } diff --git a/internal/service/dataset/ingestion.go b/internal/service/dataset/ingestion.go index 792095cf6b..84750a8c47 100644 --- a/internal/service/dataset/ingestion.go +++ b/internal/service/dataset/ingestion.go @@ -12,23 +12,23 @@ import ( func (d *DatasetService) GetIngestionSummary(ctx context.Context, datasetID, userID string) (map[string]interface{}, common.ErrorCode, error) { if datasetID == "" { - return nil, common.CodeDataError, errors.New(`Lack of "Dataset ID"`) + return nil, common.CodeDataError, errors.New(`lack of "Dataset ID"`) } if !d.kbDAO.Accessible(ctx, dao.DB, datasetID, userID) { - return nil, common.CodeDataError, errors.New("No authorization.") + return nil, common.CodeDataError, errors.New("no authorization") } kb, err := d.kbDAO.GetByID(ctx, dao.DB, datasetID) if err != nil { if dao.IsNotFoundErr(err) { - return nil, common.CodeDataError, fmt.Errorf("Invalid Dataset ID '%s'", datasetID) + return nil, common.CodeDataError, fmt.Errorf("invalid Dataset ID '%s'", datasetID) } - return nil, common.CodeServerError, errors.New("Database operation failed") + return nil, common.CodeServerError, errors.New("database operation failed") } status, err := d.documentDAO.GetParsingStatusByKBID(ctx, dao.DB, datasetID) if err != nil { - return nil, common.CodeServerError, errors.New("Database operation failed") + return nil, common.CodeServerError, errors.New("database operation failed") } return map[string]interface{}{ @@ -41,10 +41,10 @@ func (d *DatasetService) GetIngestionSummary(ctx context.Context, datasetID, use func (d *DatasetService) ListIngestionLogs(ctx context.Context, datasetID, userID string, page, pageSize int, orderby string, desc bool, operationStatus []string, createDateFrom, createDateTo, logType, keywords string) (map[string]interface{}, common.ErrorCode, error) { if datasetID == "" { - return nil, common.CodeDataError, errors.New(`Lack of "Dataset ID"`) + return nil, common.CodeDataError, errors.New(`lack of "Dataset ID"`) } if !d.kbDAO.Accessible(ctx, dao.DB, datasetID, userID) { - return nil, common.CodeDataError, errors.New("No authorization.") + return nil, common.CodeDataError, errors.New("no authorization") } if page <= 0 { @@ -63,9 +63,9 @@ func (d *DatasetService) ListIngestionLogs(ctx context.Context, datasetID, userI err error ) if logType == "file" { - logs, total, err = d.pipelineLogDAO.GetFileLogsByKBID(datasetID, page, pageSize, orderby, desc, keywords, operationStatus, createDateFrom, createDateTo) + logs, total, err = d.pipelineLogDAO.GetFileLogsByKBID(ctx, dao.DB, datasetID, page, pageSize, orderby, desc, keywords, operationStatus, createDateFrom, createDateTo) } else { - logs, total, err = d.pipelineLogDAO.GetDatasetLogsByKBID(datasetID, page, pageSize, orderby, desc, operationStatus, createDateFrom, createDateTo, keywords) + logs, total, err = d.pipelineLogDAO.GetDatasetLogsByKBID(ctx, dao.DB, datasetID, page, pageSize, orderby, desc, operationStatus, createDateFrom, createDateTo, keywords) } if err != nil { return nil, common.CodeServerError, fmt.Errorf("list ingestion logs: %w", err) @@ -91,22 +91,22 @@ func (d *DatasetService) ListIngestionLogs(ctx context.Context, datasetID, userI func (d *DatasetService) GetIngestionLog(ctx context.Context, datasetID, userID, logID string) (map[string]interface{}, common.ErrorCode, error) { if datasetID == "" { - return nil, common.CodeDataError, errors.New(`Lack of "Dataset ID"`) + return nil, common.CodeDataError, errors.New(`lack of "Dataset ID"`) } if logID == "" { - return nil, common.CodeDataError, errors.New(`Lack of "Log ID"`) + return nil, common.CodeDataError, errors.New(`lack of "Log ID"`) } if !d.kbDAO.Accessible(ctx, dao.DB, datasetID, userID) { - return nil, common.CodeDataError, errors.New("No authorization.") + return nil, common.CodeDataError, errors.New("no authorization") } - log, err := d.pipelineLogDAO.GetByIDAndKBID(logID, datasetID) + log, err := d.pipelineLogDAO.GetByIDAndKBID(ctx, dao.DB, logID, datasetID) if err != nil { + if dao.IsNotFoundErr(err) { + return nil, common.CodeDataError, errors.New("log not found") + } return nil, common.CodeServerError, fmt.Errorf("get ingestion log: %w", err) } - if log == nil { - return nil, common.CodeDataError, errors.New("Log not found") - } return datasetIngestionLogToMap(log), common.CodeSuccess, nil } diff --git a/internal/service/dataset/search.go b/internal/service/dataset/search.go index 0e1c3d361e..272a10a894 100644 --- a/internal/service/dataset/search.go +++ b/internal/service/dataset/search.go @@ -111,13 +111,22 @@ func (d *DatasetService) SearchDatasets(ctx context.Context, req *service.Search if searchID != "" { if d.searchService == nil { common.Warn("Search service is not initialized for search_id", zap.String("searchID", searchID)) - return nil, fmt.Errorf("Invalid search_id") + return nil, fmt.Errorf("invalid search_id") } - searchDetail, err := d.searchService.GetDetail(searchID) - if err != nil || searchDetail == nil || len(searchDetail) == 0 { + searchDetail, err := d.searchService.GetDetail(ctx, searchID) + if err != nil { + if ctx.Err() != nil { + return nil, ctx.Err() + } common.Warn("Invalid search_id", zap.String("searchID", searchID), zap.Error(err)) - return nil, fmt.Errorf("Invalid search_id") - } else if searchConfig, ok := searchDetail["search_config"].(map[string]interface{}); ok && searchConfig != nil { + return nil, fmt.Errorf("invalid search_id") + } + if searchDetail == nil || len(searchDetail) == 0 { + common.Warn("Invalid search_id", zap.String("searchID", searchID)) + return nil, fmt.Errorf("invalid search_id") + } + + if searchConfig, ok := searchDetail["search_config"].(map[string]interface{}); ok && searchConfig != nil { if scMetadataFilter, ok := searchConfig["meta_data_filter"].(map[string]interface{}); ok { metadataFilter = scMetadataFilter } @@ -155,7 +164,7 @@ func (d *DatasetService) SearchDatasets(ctx context.Context, req *service.Search chatID, _ = searchConfig["chat_id"].(string) } else { common.Warn("Invalid search_id: search_config missing or invalid", zap.String("searchID", searchID)) - return nil, fmt.Errorf("Invalid search_id") + return nil, fmt.Errorf("invalid search_id") } } @@ -165,7 +174,7 @@ func (d *DatasetService) SearchDatasets(ctx context.Context, req *service.Search method, _ := metadataFilter["method"].(string) if method == "auto" || method == "semi_auto" { if chatID != "" { - driver, modelName, apiConfig, _, err := modelProviderSvc.ResolveModelConfig(tenantIDs[0], entity.ModelTypeChat, chatID) + driver, modelName, apiConfig, _, err := modelProviderSvc.ResolveModelConfig(ctx, tenantIDs[0], entity.ModelTypeChat, chatID) if err != nil { common.Warn("Failed to get chat model config from search_config chat_id, using tenant default", zap.String("chatID", chatID), zap.Error(err)) } else { @@ -174,7 +183,7 @@ func (d *DatasetService) SearchDatasets(ctx context.Context, req *service.Search } if chatModelForFilter == nil { - driver, modelName, apiConfig, _, err := modelProviderSvc.GetTenantDefaultModelByType(tenantIDs[0], entity.ModelTypeChat) + driver, modelName, apiConfig, _, err := modelProviderSvc.GetTenantDefaultModelByType(ctx, tenantIDs[0], entity.ModelTypeChat) if err != nil { common.Warn("Failed to get tenant default chat model for meta_data_filter", zap.Error(err)) } else { @@ -209,7 +218,7 @@ func (d *DatasetService) SearchDatasets(ctx context.Context, req *service.Search } } if keyword { - driver, modelName, apiConfig, _, err := modelProviderSvc.GetTenantDefaultModelByType(tenantIDs[0], entity.ModelTypeChat) + driver, modelName, apiConfig, _, err := modelProviderSvc.GetTenantDefaultModelByType(ctx, tenantIDs[0], entity.ModelTypeChat) if err != nil { common.Warn("Failed to get default chat model for LLM transformations", zap.Error(err)) } else { @@ -230,7 +239,7 @@ func (d *DatasetService) SearchDatasets(ctx context.Context, req *service.Search // Determine embedding model var embeddingModel *modelModule.EmbeddingModel if kbRecords[0].EmbdID != "" { - driver, modelName, apiConfig, maxTokens, embErr := modelProviderSvc.ResolveModelConfig(tenantIDs[0], entity.ModelTypeEmbedding, kbRecords[0].EmbdID) + driver, modelName, apiConfig, maxTokens, embErr := modelProviderSvc.ResolveModelConfig(ctx, tenantIDs[0], entity.ModelTypeEmbedding, kbRecords[0].EmbdID) if embErr != nil { return nil, fmt.Errorf("failed to get embedding model by embd_id: %w", embErr) } @@ -240,7 +249,7 @@ func (d *DatasetService) SearchDatasets(ctx context.Context, req *service.Search // Get rerank model if rerankID is specified var rerankModel *modelModule.RerankModel if rerankID != "" { - driver, modelName, apiConfig, _, rErr := modelProviderSvc.ResolveModelConfig(tenantIDs[0], entity.ModelTypeRerank, rerankID) + driver, modelName, apiConfig, _, rErr := modelProviderSvc.ResolveModelConfig(ctx, tenantIDs[0], entity.ModelTypeRerank, rerankID) if rErr != nil { return nil, fmt.Errorf("failed to get rerank model by rerank_id: %w", rErr) } diff --git a/internal/service/dataset/update.go b/internal/service/dataset/update.go index d7b2d820c2..e2934f41df 100644 --- a/internal/service/dataset/update.go +++ b/internal/service/dataset/update.go @@ -212,13 +212,13 @@ func (d *DatasetService) UpdateDataset(ctx context.Context, datasetID, tenantID } else { tenantEmbdID = "" } - ok, message := d.verifyEmbeddingAvailability(effectiveEmbdID, tenantID) + ok, message := d.verifyEmbeddingAvailability(ctx, effectiveEmbdID, tenantID) if !ok { txCode = common.CodeDataError return errors.New(message) } if effectiveEmbdID != "" && tenantEmbdID == "" { - resolvedID, err := service.NewModelProviderService().ResolveModelID(tenantID, entity.ModelTypeEmbedding, effectiveEmbdID) + resolvedID, err := service.NewModelProviderService().ResolveModelID(ctx, tenantID, entity.ModelTypeEmbedding, effectiveEmbdID) if err == nil { tenantEmbdID = resolvedID } diff --git a/internal/service/generator.go b/internal/service/generator.go index f22a85122d..ae7a84df5e 100644 --- a/internal/service/generator.go +++ b/internal/service/generator.go @@ -122,13 +122,13 @@ func CrossLanguages(ctx context.Context, tenantID string, llmID string, query st break } } - driver, modelName, apiConfig, _, err := modelProviderSvc.ResolveModelConfig(tenantID, resolvedType, llmID) + driver, modelName, apiConfig, _, err := modelProviderSvc.ResolveModelConfig(ctx, tenantID, resolvedType, llmID) if err != nil { return query, fmt.Errorf("failed to get chat model: %w", err) } chatModel = modelModule.NewChatModel(driver, &modelName, apiConfig) } else { - driver, modelName, apiConfig, _, err := modelProviderSvc.GetTenantDefaultModelByType(tenantID, entity.ModelTypeChat) + driver, modelName, apiConfig, _, err := modelProviderSvc.GetTenantDefaultModelByType(ctx, tenantID, entity.ModelTypeChat) if err != nil { return query, fmt.Errorf("failed to get default chat model: %w", err) } diff --git a/internal/service/memory.go b/internal/service/memory.go index cd11b1d3e8..a11427db38 100644 --- a/internal/service/memory.go +++ b/internal/service/memory.go @@ -349,7 +349,7 @@ func (s *MemoryService) CreateMemory(ctx context.Context, tenantID string, req * // tenant_model row) — we leave the tenant_*_id fields nil and proceed. modelProvider := NewModelProviderService() if req.LLMID != "" && req.TenantLLMID == nil { - tenantLLMID, err := modelProvider.ResolveModelID(tenantID, entity.ModelTypeChat, req.LLMID) + tenantLLMID, err := modelProvider.ResolveModelID(ctx, tenantID, entity.ModelTypeChat, req.LLMID) if err != nil { slog.Warn("CreateMemory: failed to resolve tenant LLM id", "tenant_id", tenantID, "llm_id", req.LLMID, "err", err) } else if tenantLLMID != "" { @@ -357,7 +357,7 @@ func (s *MemoryService) CreateMemory(ctx context.Context, tenantID string, req * } } if req.EmbdID != "" && req.TenantEmbdID == nil { - tenantEmbdID, err := modelProvider.ResolveModelID(tenantID, entity.ModelTypeEmbedding, req.EmbdID) + tenantEmbdID, err := modelProvider.ResolveModelID(ctx, tenantID, entity.ModelTypeEmbedding, req.EmbdID) if err != nil { slog.Warn("CreateMemory: failed to resolve tenant embedding id", "tenant_id", tenantID, "embd_id", req.EmbdID, "err", err) } else if tenantEmbdID != "" { @@ -497,7 +497,7 @@ func (s *MemoryService) UpdateMemory(ctx context.Context, tenantID string, memor if req.LLMID != nil { updateDict["llm_id"] = *req.LLMID if req.TenantLLMID == nil && *req.LLMID != "" { - resolved, err := modelProvider.ResolveModelID(tenantID, entity.ModelTypeChat, *req.LLMID) + resolved, err := modelProvider.ResolveModelID(ctx, tenantID, entity.ModelTypeChat, *req.LLMID) if err != nil { slog.Warn("UpdateMemory: failed to resolve tenant LLM id", "tenant_id", tenantID, "llm_id", *req.LLMID, "err", err) } else if resolved != "" { @@ -509,7 +509,7 @@ func (s *MemoryService) UpdateMemory(ctx context.Context, tenantID string, memor if req.EmbdID != nil { updateDict["embd_id"] = *req.EmbdID if req.TenantEmbdID == nil && *req.EmbdID != "" { - resolved, err := modelProvider.ResolveModelID(tenantID, entity.ModelTypeEmbedding, *req.EmbdID) + resolved, err := modelProvider.ResolveModelID(ctx, tenantID, entity.ModelTypeEmbedding, *req.EmbdID) if err != nil { slog.Warn("UpdateMemory: failed to resolve tenant embedding id", "tenant_id", tenantID, "embd_id", *req.EmbdID, "err", err) } else if resolved != "" { @@ -1290,7 +1290,7 @@ func memoryMessageTextExpr(question string, similarityThreshold float64) *engine } func (s *MemoryService) memoryMessageDenseExpr(ctx context.Context, question string, memory *entity.Memory, topN int, similarityThreshold float64) (*enginetypes.MatchDenseExpr, error) { - driver, modelName, apiConfig, maxTokens, err := NewModelProviderService().ResolveModelConfig(memory.TenantID, entity.ModelTypeEmbedding, memory.EmbdID) + driver, modelName, apiConfig, maxTokens, err := NewModelProviderService().ResolveModelConfig(ctx, memory.TenantID, entity.ModelTypeEmbedding, memory.EmbdID) if err != nil { return nil, err } @@ -1443,7 +1443,7 @@ func (s *MemoryService) ListMemories(ctx context.Context, userID string, tenantI // If tenantIDs is empty, get all tenants associated with the user if len(tenantIDs) == 0 { userTenantService := NewUserTenantService() - userTenants, err := userTenantService.GetUserTenantRelationByUserID(userID) + userTenants, err := userTenantService.GetUserTenantRelationByUserIDWithContext(ctx, userID) if err != nil { return nil, fmt.Errorf("failed to get user tenants: %w", err) } diff --git a/internal/service/memory_message_service.go b/internal/service/memory_message_service.go index f7c9548377..da7f55e67f 100644 --- a/internal/service/memory_message_service.go +++ b/internal/service/memory_message_service.go @@ -262,7 +262,7 @@ func (s *MemoryMessageService) embedAndSave(ctx context.Context, mem *CreateMemo } content, _ := rawMessage["content"].(string) - driver, modelName, apiConfig, maxTokens, err := NewModelProviderService().ResolveModelConfig(mem.TenantID, entity.ModelTypeEmbedding, mem.EmbdID) + driver, modelName, apiConfig, maxTokens, err := NewModelProviderService().ResolveModelConfig(ctx, mem.TenantID, entity.ModelTypeEmbedding, mem.EmbdID) if err != nil { return err } diff --git a/internal/service/model_service.go b/internal/service/model_service.go index d31b26c662..207d59cec2 100644 --- a/internal/service/model_service.go +++ b/internal/service/model_service.go @@ -3214,7 +3214,7 @@ func (m *ModelProviderService) ParseFile(ctx context.Context, providerName, inst // GetEmbeddingModel returns an EmbeddingModel wrapper for the given tenant func (m *ModelProviderService) GetEmbeddingModel(ctx context.Context, tenantID, compositeModelName string) (*modelModule.EmbeddingModel, error) { - driver, modelName, apiConfig, maxTokens, err := m.ResolveModelConfig(tenantID, entity.ModelTypeEmbedding, compositeModelName) + driver, modelName, apiConfig, maxTokens, err := m.ResolveModelConfig(ctx, tenantID, entity.ModelTypeEmbedding, compositeModelName) if err != nil { return nil, err } @@ -3223,7 +3223,7 @@ func (m *ModelProviderService) GetEmbeddingModel(ctx context.Context, tenantID, // GetChatModel returns a ChatModel wrapper for the given tenant func (m *ModelProviderService) GetChatModel(ctx context.Context, tenantID, compositeModelName string) (*modelModule.ChatModel, error) { - driver, modelName, apiConfig, _, err := m.ResolveModelConfig(tenantID, entity.ModelTypeChat, compositeModelName) + driver, modelName, apiConfig, _, err := m.ResolveModelConfig(ctx, tenantID, entity.ModelTypeChat, compositeModelName) if err != nil { return nil, err } @@ -3231,8 +3231,8 @@ func (m *ModelProviderService) GetChatModel(ctx context.Context, tenantID, compo } // GetRerankModel returns a RerankModel wrapper for the given tenant -func (m *ModelProviderService) GetRerankModel(tenantID, compositeModelName string) (*modelModule.RerankModel, error) { - driver, modelName, apiConfig, _, err := m.ResolveModelConfig(tenantID, entity.ModelTypeRerank, compositeModelName) +func (m *ModelProviderService) GetRerankModel(ctx context.Context, tenantID, compositeModelName string) (*modelModule.RerankModel, error) { + driver, modelName, apiConfig, _, err := m.ResolveModelConfig(ctx, tenantID, entity.ModelTypeRerank, compositeModelName) if err != nil { return nil, err } @@ -3248,7 +3248,7 @@ type AddModelRequest struct { Extra map[string]interface{} `json:"extra"` } -func (m *ModelProviderService) GetTenantDefaultModelByType(tenantID string, modelType entity.ModelType) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { +func (m *ModelProviderService) GetTenantDefaultModelByType(ctx context.Context, tenantID string, modelType entity.ModelType) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { if modelType == entity.ModelTypeOCR { return nil, "", nil, 0, fmt.Errorf("OCR model name is required") } @@ -3259,7 +3259,7 @@ func (m *ModelProviderService) GetTenantDefaultModelByType(tenantID string, mode } modelName, modelID := defaultModelRefs(tenant, modelType) if modelID != "" { - driver, resolvedName, apiConfig, maxTokens, idErr := m.GetModelConfigByID(tenantID, modelType, modelID) + driver, resolvedName, apiConfig, maxTokens, idErr := m.GetModelConfigByID(ctx, tenantID, modelType, modelID) if idErr == nil { return driver, resolvedName, apiConfig, maxTokens, nil } @@ -3272,11 +3272,11 @@ func (m *ModelProviderService) GetTenantDefaultModelByType(tenantID string, mode return nil, "", nil, 0, fmt.Errorf("no default %s model is set", modelType) } - return m.ResolveModelConfig(tenantID, modelType, modelName) + return m.ResolveModelConfig(ctx, tenantID, modelType, modelName) } // GetModelConfigByID returns model driver and API config for a tenant_model row by its ID. -func (m *ModelProviderService) GetModelConfigByID(userID string, modelType entity.ModelType, modelID string) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { +func (m *ModelProviderService) GetModelConfigByID(ctx context.Context, userID string, modelType entity.ModelType, modelID string) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { common.Debug("GetModelConfigByID", zap.String("userID", userID), zap.String("modelType", modelType.String()), @@ -3308,7 +3308,7 @@ func (m *ModelProviderService) GetModelConfigByID(userID string, modelType entit } if providerEntity.TenantID != userID { - userTenants, terr := NewUserTenantService().GetUserTenantRelationByUserID(userID) + userTenants, terr := NewUserTenantService().GetUserTenantRelationByUserIDWithContext(ctx, userID) if terr != nil { return nil, "", nil, 0, terr } @@ -3388,19 +3388,19 @@ func defaultModelRefs(tenant *entity.Tenant, modelType entity.ModelType) (string } } -func (m *ModelProviderService) ResolveModelConfig(tenantID string, modelType entity.ModelType, modelRef string) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { +func (m *ModelProviderService) ResolveModelConfig(ctx context.Context, tenantID string, modelType entity.ModelType, modelRef string) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { if strings.TrimSpace(modelRef) == "" { return nil, "", nil, 0, fmt.Errorf("model ref is required") } if _, err := m.modelDAO.GetByID(modelRef); err == nil { - return m.GetModelConfigByID(tenantID, modelType, modelRef) + return m.GetModelConfigByID(ctx, tenantID, modelType, modelRef) } else if !errors.Is(err, gorm.ErrRecordNotFound) { return nil, "", nil, 0, err } return m.GetModelConfigFromProviderInstance(tenantID, modelType, modelRef) } -func (m *ModelProviderService) ResolveModelID(tenantID string, modelType entity.ModelType, modelName string) (string, error) { +func (m *ModelProviderService) ResolveModelID(ctx context.Context, tenantID string, modelType entity.ModelType, modelName string) (string, error) { if modelObj, err := m.modelDAO.GetByID(modelName); err == nil { if modelObj.Status != "active" { return "", fmt.Errorf("tenant model id=%s is disabled", modelName) @@ -3408,7 +3408,7 @@ func (m *ModelProviderService) ResolveModelID(tenantID string, modelType entity. if !entity.ModelType(modelObj.ModelType).Has(modelType) { return "", fmt.Errorf("tenant model id=%s cannot be used as %s model", modelName, modelType.String()) } - if _, _, _, _, err := m.GetModelConfigByID(tenantID, modelType, modelName); err != nil { + if _, _, _, _, err = m.GetModelConfigByID(ctx, tenantID, modelType, modelName); err != nil { return "", err } return modelObj.ID, nil @@ -4026,13 +4026,13 @@ func (m *ModelProviderService) isImage2TextLLM(tenantID, llmID string) bool { // If llmID is empty, falls back to the tenant's default chat model. // When the named LLM is registered as an image2text model, returns the // IMAGE2TEXT driver/config instead of CHAT. -func (m *ModelProviderService) GetChatModelConfig(tenantID string, llmID string) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { +func (m *ModelProviderService) GetChatModelConfig(ctx context.Context, tenantID string, llmID string) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { if llmID == "" { - return m.GetTenantDefaultModelByType(tenantID, entity.ModelTypeChat) + return m.GetTenantDefaultModelByType(ctx, tenantID, entity.ModelTypeChat) } modelType := entity.ModelTypeChat if m.isImage2TextLLM(tenantID, llmID) { modelType = entity.ModelTypeImage2Text } - return m.ResolveModelConfig(tenantID, modelType, llmID) + return m.ResolveModelConfig(ctx, tenantID, modelType, llmID) } diff --git a/internal/service/model_service_test.go b/internal/service/model_service_test.go index 0f031667e4..3c163ce547 100644 --- a/internal/service/model_service_test.go +++ b/internal/service/model_service_test.go @@ -182,7 +182,8 @@ func TestModelProviderServiceGetModelConfigByID(t *testing.T) { useModelProviderServiceTestDB(t, db) seedModelProviderServiceScope(t, db) - driver, modelName, apiConfig, _, err := NewModelProviderService().GetModelConfigByID("user-1", entity.ModelTypeChat, "model-1") + ctx := t.Context() + driver, modelName, apiConfig, _, err := NewModelProviderService().GetModelConfigByID(ctx, "user-1", entity.ModelTypeChat, "model-1") if err != nil { t.Fatalf("GetModelConfigByID() error = %v", err) } diff --git a/internal/service/openai_chat.go b/internal/service/openai_chat.go index ef98986187..9d9347a2ac 100644 --- a/internal/service/openai_chat.go +++ b/internal/service/openai_chat.go @@ -232,7 +232,7 @@ func (s *OpenAIChatService) OpenAIChatCompletions(c *gin.Context, userID, chatID } } if req.Model != "model" { - if _, _, _, _, mErr := s.pipeline.ModelProviderSvc.GetChatModelConfig(dialog.TenantID, resolvedModel); mErr != nil { + if _, _, _, _, mErr := s.pipeline.ModelProviderSvc.GetChatModelConfig(ctx, dialog.TenantID, resolvedModel); mErr != nil { s.writeArgError(c, fmt.Sprintf("`llm_id` %s doesn't exist", req.Model)) return } diff --git a/internal/service/related_question.go b/internal/service/related_question.go index a582eb810e..a64adfef00 100644 --- a/internal/service/related_question.go +++ b/internal/service/related_question.go @@ -32,7 +32,7 @@ func GenerateRelatedQuestions(ctx context.Context, tenantID, question, searchID if modelProviderSvc == nil { return nil, fmt.Errorf("model provider service not configured") } - searchConfig := relatedQuestionsSearchConfig(searchID, searchSvc) + searchConfig := relatedQuestionsSearchConfig(ctx, searchID, searchSvc) modelID := relatedQuestionsModelID(tenantID, searchConfig, tenantSvc) prompt, err := LoadPrompt("related_question") if err != nil { @@ -52,11 +52,11 @@ func GenerateRelatedQuestions(ctx context.Context, tenantID, question, searchID return []string{}, nil } -func relatedQuestionsSearchConfig(searchID string, searchSvc *SearchService) map[string]interface{} { +func relatedQuestionsSearchConfig(ctx context.Context, searchID string, searchSvc *SearchService) map[string]interface{} { if searchID == "" || searchSvc == nil { return map[string]interface{}{} } - if detail, err := searchSvc.GetDetail(searchID); err == nil && detail != nil { + if detail, err := searchSvc.GetDetail(ctx, searchID); err == nil && detail != nil { return relatedQuestionsSearchConfigFromDetail(detail) } return map[string]interface{}{} diff --git a/internal/service/search.go b/internal/service/search.go index 797706f968..4d75392c78 100644 --- a/internal/service/search.go +++ b/internal/service/search.go @@ -80,13 +80,13 @@ type SearchShareDetail struct { } // ListSearches list search apps with advanced filtering (equivalent to list_search_app) -func (s *SearchService) ListSearches(userID string, keywords string, page, pageSize int, orderby string, desc bool, ownerIDs []string) (*ListSearchAppsResponse, error) { +func (s *SearchService) ListSearches(ctx context.Context, userID string, keywords string, page, pageSize int, orderby string, desc bool, ownerIDs []string) (*ListSearchAppsResponse, error) { var searches []*entity.SearchListItem var total int64 var err error if len(ownerIDs) == 0 { - searches, total, err = s.searchDAO.ListByTenantIDs(nil, userID, page, pageSize, orderby, desc, keywords) + searches, total, err = s.searchDAO.ListByTenantIDs(ctx, dao.DB, nil, userID, page, pageSize, orderby, desc, keywords) if err != nil { return nil, err } @@ -102,7 +102,7 @@ func (s *SearchService) ListSearches(userID string, keywords string, page, pageS }, nil } - searches, total, err = s.searchDAO.ListByOwnerIDs(ownerIDs, userID, orderby, desc, keywords) + searches, total, err = s.searchDAO.ListByOwnerIDs(ctx, dao.DB, ownerIDs, userID, orderby, desc, keywords) if err != nil { return nil, err } @@ -207,7 +207,7 @@ type CreateSearchResponse struct { // 5. Set fields: id, name, description, tenant_id, created_by // 6. Save to database within DB.atomic() transaction // 7. Return {search_id: id} on success -func (s *SearchService) CreateSearch(userID string, name string, description *string) (*CreateSearchResponse, error) { +func (s *SearchService) CreateSearch(ctx context.Context, userID string, name string, description *string) (*CreateSearchResponse, error) { if err := common.ValidateName(name); err != nil { return nil, err } @@ -217,7 +217,7 @@ func (s *SearchService) CreateSearch(userID string, name string, description *st // Generate unique name (same as Python duplicate_name) uniqueName, err := common.DuplicateName(func(name string, tid string) bool { - existing, _ := s.searchDAO.GetByNameAndTenant(name, tid) + existing, _ := s.searchDAO.GetByNameAndTenant(ctx, dao.DB, name, tid) return len(existing) > 0 }, name, userID) @@ -243,7 +243,7 @@ func (s *SearchService) CreateSearch(userID string, name string, description *st search.Status = &status // Save to database - if err := s.searchDAO.Create(search); err != nil { + if err = s.searchDAO.Create(ctx, dao.DB, search); err != nil { return nil, fmt.Errorf("failed to create search: %w", err) } @@ -252,7 +252,7 @@ func (s *SearchService) CreateSearch(userID string, name string, description *st }, nil } -func (s *SearchService) GetSearchDetail(userID string, searchID string) (*entity.Search, error) { +func (s *SearchService) GetSearchDetail(ctx context.Context, userID string, searchID string) (*entity.Search, error) { // Step 1: Get user tenants (same as Python UserTenantService.query(user_id=current_user.id)) tenants, err := s.userTenantDAO.GetByUserID(userID) if err != nil { @@ -263,8 +263,11 @@ func (s *SearchService) GetSearchDetail(userID string, searchID string) (*entity // Python: for tenant in tenants: if SearchService.query(tenant_id=tenant.tenant_id, id=search_id): break hasPermission := false for _, tenant := range tenants { - searches, err := s.searchDAO.QueryByTenantIDAndID(tenant.TenantID, searchID) + searches, err := s.searchDAO.QueryByTenantIDAndID(ctx, dao.DB, tenant.TenantID, searchID) if err != nil { + if ctx.Err() != nil { + return nil, ctx.Err() + } continue // Try next tenant } if len(searches) > 0 { @@ -278,7 +281,7 @@ func (s *SearchService) GetSearchDetail(userID string, searchID string) (*entity } // Step 3: Get search detail (same as Python SearchService.get_detail(search_id)) - search, err := s.searchDAO.GetByID(searchID) + search, err := s.searchDAO.GetByID(ctx, dao.DB, searchID) if err != nil { return nil, fmt.Errorf("can't find this Search App!") } @@ -288,12 +291,12 @@ func (s *SearchService) GetSearchDetail(userID string, searchID string) (*entity // GetSearchShareDetail returns the joined share-detail payload for public // searchbot pages after verifying the caller can access the search app. -func (s *SearchService) GetSearchShareDetail(userID, searchID string) (*SearchShareDetail, error) { - if _, err := s.GetSearchDetail(userID, searchID); err != nil { +func (s *SearchService) GetSearchShareDetail(ctx context.Context, userID, searchID string) (*SearchShareDetail, error) { + if _, err := s.GetSearchDetail(ctx, userID, searchID); err != nil { return nil, err } - detail, err := s.searchDAO.GetDetailByID(searchID) + detail, err := s.searchDAO.GetDetailByID(ctx, dao.DB, searchID) if err != nil { return nil, err } @@ -314,11 +317,11 @@ func (s *SearchService) GetSearchShareDetail(userID, searchID string) (*SearchSh } // DeleteSearch deletes a search app by ID -func (s *SearchService) DeleteSearch(userID string, searchID string) error { +func (s *SearchService) DeleteSearch(ctx context.Context, userID string, searchID string) error { // Step 1: Check deletion permission (same as Python SearchService.accessible4deletion) // Python: cls.model.select().where(cls.model.id == search_id, cls.model.created_by == user_id, cls.model.status == StatusEnum.VALID.value).first() - status, err := s.searchDAO.Accessible4Deletion(searchID, userID) + status, err := s.searchDAO.Accessible4Deletion(ctx, dao.DB, searchID, userID) if err != nil { return fmt.Errorf("failed to check deletion permission: %w", err) } @@ -329,7 +332,7 @@ func (s *SearchService) DeleteSearch(userID string, searchID string) error { // Step 2: Execute delete (same as Python SearchService.delete_by_id) // Python: cls.model.delete().where(cls.model.id == pid).execute() - if err = s.searchDAO.DeleteByID(userID, searchID); err != nil { + if err = s.searchDAO.DeleteByID(ctx, dao.DB, userID, searchID); err != nil { return fmt.Errorf("failed to delete search App %s: %w", searchID, err) } @@ -337,8 +340,8 @@ func (s *SearchService) DeleteSearch(userID string, searchID string) error { } // AccessibleForCompletion check if it is accessible -func (s *SearchService) AccessibleForCompletion(userID string, searchID string) (bool, error) { - return s.searchDAO.Accessible4Deletion(searchID, userID) +func (s *SearchService) AccessibleForCompletion(ctx context.Context, userID string, searchID string) (bool, error) { + return s.searchDAO.Accessible4Deletion(ctx, dao.DB, searchID, userID) } type SearchCompletionPlan struct { @@ -367,7 +370,7 @@ func (s *SearchService) PrepareCompletion(ctx context.Context, userID, searchID return nil, common.CodeArgumentError, fmt.Errorf("question is required") } - accessible, err := s.AccessibleForCompletion(userID, searchID) + accessible, err := s.AccessibleForCompletion(ctx, userID, searchID) if err != nil { return nil, common.CodeServerError, err } @@ -375,7 +378,7 @@ func (s *SearchService) PrepareCompletion(ctx context.Context, userID, searchID return nil, common.CodeAuthenticationError, fmt.Errorf("no authorization") } - searchDetail, err := s.GetDetail(searchID) + searchDetail, err := s.GetDetail(ctx, searchID) if err != nil || searchDetail == nil { return nil, common.CodeDataError, fmt.Errorf("cannot find search %s", searchID) } @@ -466,7 +469,7 @@ func searchConfigMapValue(value interface{}) (map[string]interface{}, bool) { case map[string]interface{}: return typed, true case entity.JSONMap: - return map[string]interface{}(typed), true + return typed, true default: return nil, false } @@ -553,12 +556,12 @@ type UpdateSearchRequest struct { Avatar *string `json:"avatar,omitempty"` } -func (s *SearchService) UpdateSearch(userID string, searchID string, req *UpdateSearchRequest) (*entity.Search, error) { +func (s *SearchService) UpdateSearch(ctx context.Context, userID string, searchID string, req *UpdateSearchRequest) (*entity.Search, error) { // Step 1: Check update permission (same as delete - uses accessible4deletion) // Only creator can update. A missing or non-owned search is treated as // unauthorized so the contract returns a clear "no authorization" error. - accessible, err := s.searchDAO.Accessible4Deletion(searchID, userID) + accessible, err := s.searchDAO.Accessible4Deletion(ctx, dao.DB, searchID, userID) if err != nil { return nil, fmt.Errorf("failed to check deletion permission: %w", err) } @@ -568,7 +571,7 @@ func (s *SearchService) UpdateSearch(userID string, searchID string, req *Update // Step 2: Get existing search // Python: search_app = SearchService.query(tenant_id=current_user.id, id=search_id)[0] - search, err := s.searchDAO.GetByTenantIDAndID(userID, searchID) + search, err := s.searchDAO.GetByTenantIDAndID(ctx, dao.DB, userID, searchID) if err != nil { return nil, fmt.Errorf("cannot find search %s", searchID) } @@ -577,7 +580,7 @@ func (s *SearchService) UpdateSearch(userID string, searchID string, req *Update // Python: if req["name"].lower() != search_app.name.lower() and len(SearchService.query(...)) >= 1 trimmedName := req.Name if search.Name != trimmedName { - existing, _ := s.searchDAO.GetByNameAndTenant(trimmedName, userID) + existing, _ := s.searchDAO.GetByNameAndTenant(ctx, dao.DB, trimmedName, userID) if len(existing) > 0 { return nil, fmt.Errorf("duplicated search name") } @@ -615,13 +618,13 @@ func (s *SearchService) UpdateSearch(userID string, searchID string, req *Update // Step 6: Execute update // Python: SearchService.update_by_id(search_id, req) - if err = s.searchDAO.UpdateByID(searchID, updates); err != nil { + if err = s.searchDAO.UpdateByID(ctx, dao.DB, searchID, updates); err != nil { return nil, fmt.Errorf("failed to update search: %w", err) } // Step 7: Fetch updated search // Python: e, updated_search = SearchService.get_by_id(search_id) - updatedSearch, err := s.searchDAO.GetByID(searchID) + updatedSearch, err := s.searchDAO.GetByID(ctx, dao.DB, searchID) if err != nil { return nil, fmt.Errorf("failed to fetch updated search: %w", err) } @@ -630,8 +633,8 @@ func (s *SearchService) UpdateSearch(userID string, searchID string, req *Update } // GetDetail gets search details by ID including search_config -func (s *SearchService) GetDetail(searchID string) (map[string]interface{}, error) { - search, err := s.searchDAO.GetByID(searchID) +func (s *SearchService) GetDetail(ctx context.Context, searchID string) (map[string]interface{}, error) { + search, err := s.searchDAO.GetByID(ctx, dao.DB, searchID) if err != nil { return nil, err diff --git a/internal/service/search_share_detail_test.go b/internal/service/search_share_detail_test.go index 87ec220abe..cf61842e54 100644 --- a/internal/service/search_share_detail_test.go +++ b/internal/service/search_share_detail_test.go @@ -77,7 +77,8 @@ func TestSearchServiceGetSearchShareDetail(t *testing.T) { t.Fatalf("failed to create search: %v", err) } - detail, err := NewSearchService().GetSearchShareDetail("tenant-1", "search-1") + ctx := t.Context() + detail, err := NewSearchService().GetSearchShareDetail(ctx, "tenant-1", "search-1") if err != nil { t.Fatalf("GetSearchShareDetail failed: %v", err) } @@ -144,7 +145,8 @@ func TestSearchServiceGetSearchShareDetailRejectsUnauthorizedUser(t *testing.T) t.Fatalf("failed to create search: %v", err) } - _, err := NewSearchService().GetSearchShareDetail("user-2", "search-1") + ctx := t.Context() + _, err := NewSearchService().GetSearchShareDetail(ctx, "user-2", "search-1") if err == nil { t.Fatal("expected permission error") } diff --git a/internal/service/search_test.go b/internal/service/search_test.go index a117ca8267..31b7608a29 100644 --- a/internal/service/search_test.go +++ b/internal/service/search_test.go @@ -52,7 +52,8 @@ func createSearchServiceTestSearch(t *testing.T, id, tenantID, name string) { func TestSearchServiceCreateRejectsEmptyName(t *testing.T) { setupSearchServiceTestDB(t) - _, err := NewSearchService().CreateSearch("tenant-1", " ", nil) + ctx := t.Context() + _, err := NewSearchService().CreateSearch(ctx, "tenant-1", " ", nil) if err == nil { t.Fatal("expected empty name validation error") } @@ -64,11 +65,12 @@ func TestSearchServiceCreateRejectsEmptyName(t *testing.T) { func TestSearchServiceUpdateRejectsUnauthorizedSearchID(t *testing.T) { setupSearchServiceTestDB(t) + ctx := t.Context() req := &UpdateSearchRequest{ Name: "New Name", SearchConfig: map[string]interface{}{}, } - _, err := NewSearchService().UpdateSearch("user-2", "invalid_search_id", req) + _, err := NewSearchService().UpdateSearch(ctx, "user-2", "invalid_search_id", req) if err == nil { t.Fatal("expected authorization error") } @@ -79,8 +81,8 @@ func TestSearchServiceUpdateRejectsUnauthorizedSearchID(t *testing.T) { func TestSearchServiceCreateAndUpdateRoundTrip(t *testing.T) { setupSearchServiceTestDB(t) - - created, err := NewSearchService().CreateSearch("tenant-1", "My Search", nil) + ctx := t.Context() + created, err := NewSearchService().CreateSearch(ctx, "tenant-1", "My Search", nil) if err != nil { t.Fatalf("CreateSearch failed: %v", err) } @@ -93,7 +95,7 @@ func TestSearchServiceCreateAndUpdateRoundTrip(t *testing.T) { Name: "Hijacked Name", SearchConfig: map[string]interface{}{}, } - _, err = NewSearchService().UpdateSearch("user-2", created.SearchID, req) + _, err = NewSearchService().UpdateSearch(ctx, "user-2", created.SearchID, req) if err == nil || err.Error() != "no authorization" { t.Fatalf("expected no authorization, got %v", err) } @@ -103,7 +105,7 @@ func TestSearchServiceCreateAndUpdateRoundTrip(t *testing.T) { Name: "Updated Name", SearchConfig: map[string]interface{}{"summary": true}, } - updated, err := NewSearchService().UpdateSearch("tenant-1", created.SearchID, req) + updated, err := NewSearchService().UpdateSearch(ctx, "tenant-1", created.SearchID, req) if err != nil { t.Fatalf("owner UpdateSearch failed: %v", err) } @@ -114,7 +116,7 @@ func TestSearchServiceCreateAndUpdateRoundTrip(t *testing.T) { t.Fatalf("expected merged search_config, got %#v", updated.SearchConfig) } - persisted, err := dao.NewSearchDAO().GetByID(created.SearchID) + persisted, err := dao.NewSearchDAO().GetByID(ctx, dao.DB, created.SearchID) if err != nil { t.Fatalf("get updated search: %v", err) } @@ -135,7 +137,8 @@ func TestSearchServiceListSearchesReturnsOwnerDisplayFields(t *testing.T) { } createSearchServiceTestSearch(t, "search-1", "user-1", "Search One") - result, err := NewSearchService().ListSearches("user-1", "", 0, 0, "create_time", true, nil) + ctx := t.Context() + result, err := NewSearchService().ListSearches(ctx, "user-1", "", 0, 0, "create_time", true, nil) if err != nil { t.Fatalf("ListSearches failed: %v", err) } @@ -152,7 +155,8 @@ func TestSearchServiceListSearchesNicknameFallsBackToTenantID(t *testing.T) { createSearchServiceTestSearch(t, "search-1", "user-1", "Search One") - result, err := NewSearchService().ListSearches("user-1", "", 0, 0, "create_time", true, nil) + ctx := t.Context() + result, err := NewSearchService().ListSearches(ctx, "user-1", "", 0, 0, "create_time", true, nil) if err != nil { t.Fatalf("ListSearches failed: %v", err) } diff --git a/internal/service/skill_indexer.go b/internal/service/skill_indexer.go index 8bf3c4ce6e..8e26427a20 100644 --- a/internal/service/skill_indexer.go +++ b/internal/service/skill_indexer.go @@ -82,7 +82,7 @@ func isElasticsearch(docEngine engine.DocEngine) bool { func (s *SkillIndexerService) IndexSkill(ctx context.Context, tenantID, spaceID string, skill SkillInfo, docEngine engine.DocEngine, embdID string) error { spaceID = normalizeSpaceID(spaceID) - config, err := s.configDAO.GetOrCreate(tenantID, spaceID, embdID) + config, err := s.configDAO.GetOrCreate(ctx, dao.DB, tenantID, spaceID, embdID) if err != nil { return fmt.Errorf("failed to get config: %w", err) } @@ -208,7 +208,7 @@ func (s *SkillIndexerService) BatchIndexSkills(ctx context.Context, tenantID, sp return nil } - config, err := s.configDAO.GetOrCreate(tenantID, spaceID, embdID) + config, err := s.configDAO.GetOrCreate(ctx, dao.DB, tenantID, spaceID, embdID) if err != nil { return fmt.Errorf("failed to get config: %w", err) } @@ -414,14 +414,14 @@ func (s *SkillIndexerService) UpdateSkillVersion(ctx context.Context, tenantID, func (s *SkillIndexerService) ReindexAll(ctx context.Context, tenantID, spaceID string, docEngine engine.DocEngine, embdID string) (map[string]interface{}, error) { spaceID = normalizeSpaceID(spaceID) // Get current config and increment semantic version - config, err := s.configDAO.GetOrCreate(tenantID, spaceID, embdID) + config, err := s.configDAO.GetOrCreate(ctx, dao.DB, tenantID, spaceID, embdID) if err != nil { return nil, fmt.Errorf("failed to get config: %w", err) } // Increment semantic version (e.g., "1.0.0" -> "1.0.1" or "1.0.9" -> "1.1.0") newVersion := incrementSemanticVersion(config.IndexVersion) - if err := s.configDAO.UpdateByTenantID(tenantID, spaceID, map[string]interface{}{ + if err = s.configDAO.UpdateByTenantID(ctx, dao.DB, tenantID, spaceID, map[string]interface{}{ "index_version": newVersion, }); err != nil { return nil, fmt.Errorf("failed to update version: %w", err) @@ -451,7 +451,7 @@ func (s *SkillIndexerService) ReindexAll(ctx context.Context, tenantID, spaceID } // Get space info to find folder ID - space, err := s.spaceDAO.GetByID(spaceID) + space, err := s.spaceDAO.GetByID(ctx, dao.DB, spaceID) if err != nil { return nil, fmt.Errorf("failed to get space: %w", err) } diff --git a/internal/service/skill_search.go b/internal/service/skill_search.go index d5c071a4a5..6f3409b129 100644 --- a/internal/service/skill_search.go +++ b/internal/service/skill_search.go @@ -60,7 +60,7 @@ type GetConfigRequest struct { } // GetConfig retrieves the search configuration for a tenant -func (s *SkillSearchService) GetConfig(tenantID, spaceID, embdID string) (map[string]interface{}, common.ErrorCode, error) { +func (s *SkillSearchService) GetConfig(ctx context.Context, tenantID, spaceID, embdID string) (map[string]interface{}, common.ErrorCode, error) { spaceID = normalizeSpaceID(spaceID) var config *entity.SkillSearchConfig var err error @@ -68,7 +68,7 @@ func (s *SkillSearchService) GetConfig(tenantID, spaceID, embdID string) (map[st if embdID == "" { // If embd_id is not provided, get the latest config for the tenant // Prioritize configs with non-empty embd_id (user-saved configs) - config, err = s.configDAO.GetLatestByTenantID(tenantID, spaceID) + config, err = s.configDAO.GetLatestByTenantID(ctx, dao.DB, tenantID, spaceID) if err != nil { // No config found, return default config config = &entity.SkillSearchConfig{ @@ -87,10 +87,10 @@ func (s *SkillSearchService) GetConfig(tenantID, spaceID, embdID string) (map[st } } } else { - config, err = s.configDAO.GetByTenantAndEmbdID(tenantID, spaceID, embdID) + config, err = s.configDAO.GetByTenantAndEmbdID(ctx, dao.DB, tenantID, spaceID, embdID) if err != nil { // Config not found, create default one - config, err = s.configDAO.GetOrCreate(tenantID, spaceID, embdID) + config, err = s.configDAO.GetOrCreate(ctx, dao.DB, tenantID, spaceID, embdID) if err != nil { return nil, common.CodeOperatingError, fmt.Errorf("failed to get or create config: %w", err) } @@ -113,7 +113,7 @@ type UpdateConfigRequest struct { } // UpdateConfig updates the search configuration for a tenant -func (s *SkillSearchService) UpdateConfig(req *UpdateConfigRequest) (map[string]interface{}, common.ErrorCode, error) { +func (s *SkillSearchService) UpdateConfig(ctx context.Context, req *UpdateConfigRequest) (map[string]interface{}, common.ErrorCode, error) { req.SpaceID = normalizeSpaceID(req.SpaceID) // Validate vector_similarity_weight if req.VectorSimilarityWeight < 0 || req.VectorSimilarityWeight > 1 { @@ -132,17 +132,17 @@ func (s *SkillSearchService) UpdateConfig(req *UpdateConfigRequest) (map[string] // Get or create config for this tenant+space (regardless of embd_id) // Each tenant+space should have only ONE config, switching embd_id updates the existing config - config, err := s.configDAO.GetLatestByTenantID(req.TenantID, req.SpaceID) + config, err := s.configDAO.GetLatestByTenantID(ctx, dao.DB, req.TenantID, req.SpaceID) if err != nil { // No config exists, create a new one - config, err = s.configDAO.CreateWithTenantSpace(req.TenantID, req.SpaceID, req.EmbdID) + config, err = s.configDAO.CreateWithTenantSpace(ctx, dao.DB, req.TenantID, req.SpaceID, req.EmbdID) if err != nil { return nil, common.CodeOperatingError, fmt.Errorf("failed to create config: %w", err) } } else { // Config exists, clean up any other active records for this tenant+space // to ensure only one active config per tenant+space - if err := s.configDAO.DeleteAllByTenantSpaceExceptID(req.TenantID, req.SpaceID, config.ID); err != nil { + if err := s.configDAO.DeleteAllByTenantSpaceExceptID(ctx, dao.DB, req.TenantID, req.SpaceID, config.ID); err != nil { common.Warn("Failed to clean up duplicate configs", zap.Error(err)) } } @@ -179,12 +179,12 @@ func (s *SkillSearchService) UpdateConfig(req *UpdateConfigRequest) (map[string] } // Update by config ID to ensure we update the correct record - if err := s.configDAO.Update(config.ID, updates); err != nil { + if err = s.configDAO.Update(ctx, dao.DB, config.ID, updates); err != nil { return nil, common.CodeOperatingError, fmt.Errorf("failed to update config: %w", err) } // Refresh config - config, err = s.configDAO.GetByID(config.ID) + config, err = s.configDAO.GetByID(ctx, dao.DB, config.ID) if err != nil { return nil, common.CodeOperatingError, fmt.Errorf("failed to refresh config: %w", err) } @@ -245,7 +245,7 @@ func (s *SkillSearchService) Search(ctx context.Context, req *SearchRequest, doc // Get config for search strategy // Use GetLatestByTenantID to prioritize configs with non-empty embd_id - config, err := s.configDAO.GetLatestByTenantID(req.TenantID, req.SpaceID) + config, err := s.configDAO.GetLatestByTenantID(ctx, dao.DB, req.TenantID, req.SpaceID) if err != nil { // Use default config if not found config = &entity.SkillSearchConfig{ diff --git a/internal/service/skill_space.go b/internal/service/skill_space.go index 09bbfaaa08..f075cf2264 100644 --- a/internal/service/skill_space.go +++ b/internal/service/skill_space.go @@ -165,7 +165,7 @@ func (s *SkillSpaceService) CreateSpace(ctx context.Context, req *CreateSpaceReq }() // Double-check after acquiring lock: Check if space with same name already exists (active status) - existingSpace, err := s.spaceDAO.GetByTenantAndName(req.TenantID, req.Name) + existingSpace, err := s.spaceDAO.GetByTenantAndName(ctx, dao.DB, req.TenantID, req.Name) if err != nil { // Space doesn't exist, continue } else if existingSpace != nil { @@ -173,7 +173,7 @@ func (s *SkillSpaceService) CreateSpace(ctx context.Context, req *CreateSpaceReq } // Check if there's a space with the same name that is currently being deleted - existingSpaceAny, err := s.spaceDAO.GetByTenantAndNameAnyStatus(req.TenantID, req.Name) + existingSpaceAny, err := s.spaceDAO.GetByTenantAndNameAnyStatus(ctx, dao.DB, req.TenantID, req.Name) if err == nil && existingSpaceAny != nil && existingSpaceAny.Status == entity.SpaceStatusDeleting { return nil, common.CodeDataError, fmt.Errorf("space with name '%s' is being deleted, please try again later", req.Name) } @@ -181,7 +181,7 @@ func (s *SkillSpaceService) CreateSpace(ctx context.Context, req *CreateSpaceReq // Check if there's a deleted/non-active space with the same name and permanently delete it // This handles the case where a previous creation failed partially // Only delete non-active spaces (status != '1') to prevent TOCTOU race - if err = s.spaceDAO.DeletePermanentByName(req.TenantID, req.Name); err != nil { + if err = s.spaceDAO.DeletePermanentByName(ctx, dao.DB, req.TenantID, req.Name); err != nil { common.Warn("Failed to delete permanent space by name", zap.Error(err)) } @@ -242,7 +242,7 @@ func (s *SkillSpaceService) CreateSpace(ctx context.Context, req *CreateSpaceReq Status: "1", } - if err = s.spaceDAO.Create(space); err != nil { + if err = s.spaceDAO.Create(ctx, dao.DB, space); err != nil { // Rollback: delete the created folder common.Error("Failed to create space in database", err) s.fileDAO.DeleteByIDs(ctx, dao.DB, []string{folderID}) @@ -261,7 +261,7 @@ func (s *SkillSpaceService) CreateSpace(ctx context.Context, req *CreateSpaceReq } } if defaultEmbdID != "" { - if _, err := s.configDAO.GetOrCreate(req.TenantID, spaceID, defaultEmbdID); err != nil { + if _, err := s.configDAO.GetOrCreate(ctx, dao.DB, req.TenantID, spaceID, defaultEmbdID); err != nil { common.Warn("Failed to create skill search config for new space", zap.String("tenantID", req.TenantID), zap.String("spaceID", spaceID), @@ -274,8 +274,8 @@ func (s *SkillSpaceService) CreateSpace(ctx context.Context, req *CreateSpaceReq } // ListSpaces lists all skills spaces for a tenant -func (s *SkillSpaceService) ListSpaces(tenantID string) (map[string]interface{}, common.ErrorCode, error) { - spaces, err := s.spaceDAO.GetByTenantID(tenantID) +func (s *SkillSpaceService) ListSpaces(ctx context.Context, tenantID string) (map[string]interface{}, common.ErrorCode, error) { + spaces, err := s.spaceDAO.GetByTenantID(ctx, dao.DB, tenantID) if err != nil { return nil, common.CodeOperatingError, fmt.Errorf("failed to list spaces: %w", err) } @@ -293,8 +293,8 @@ func (s *SkillSpaceService) ListSpaces(tenantID string) (map[string]interface{}, } // GetSpace retrieves a skills space by ID (includes deleting status for visibility) -func (s *SkillSpaceService) GetSpace(spaceID, tenantID string) (map[string]interface{}, common.ErrorCode, error) { - space, err := s.spaceDAO.GetByIDAnyStatus(spaceID) +func (s *SkillSpaceService) GetSpace(ctx context.Context, spaceID, tenantID string) (map[string]interface{}, common.ErrorCode, error) { + space, err := s.spaceDAO.GetByIDAnyStatus(ctx, dao.DB, spaceID) if err != nil { return nil, common.CodeDataError, fmt.Errorf("space not found") } @@ -314,7 +314,7 @@ func (s *SkillSpaceService) GetSpace(spaceID, tenantID string) (map[string]inter // UpdateSpace updates a skills space func (s *SkillSpaceService) UpdateSpace(ctx context.Context, spaceID string, tenantID string, req *UpdateSpaceRequest) (map[string]interface{}, common.ErrorCode, error) { - space, err := s.spaceDAO.GetByID(spaceID) + space, err := s.spaceDAO.GetByID(ctx, dao.DB, spaceID) if err != nil { return nil, common.CodeDataError, fmt.Errorf("space not found") } @@ -329,7 +329,7 @@ func (s *SkillSpaceService) UpdateSpace(ctx context.Context, spaceID string, ten if req.Name != "" && req.Name != space.Name { // Check if name already exists - existingSpace, _ := s.spaceDAO.GetByTenantAndName(tenantID, req.Name) + existingSpace, _ := s.spaceDAO.GetByTenantAndName(ctx, dao.DB, tenantID, req.Name) if existingSpace != nil && existingSpace.ID != spaceID { return nil, common.CodeDataError, fmt.Errorf("space with name '%s' already exists", req.Name) } @@ -338,7 +338,7 @@ func (s *SkillSpaceService) UpdateSpace(ctx context.Context, spaceID string, ten updates["name"] = req.Name // Update space first, then folder (atomic-like behavior with rollback on failure) - if err = s.spaceDAO.UpdateByID(spaceID, updates); err != nil { + if err = s.spaceDAO.UpdateByID(ctx, dao.DB, spaceID, updates); err != nil { return nil, common.CodeOperatingError, fmt.Errorf("failed to update space name: %w", err) } @@ -346,7 +346,7 @@ func (s *SkillSpaceService) UpdateSpace(ctx context.Context, spaceID string, ten if err = s.fileDAO.UpdateByID(ctx, dao.DB, space.FolderID, map[string]interface{}{"name": req.Name}); err != nil { common.Error("Failed to update folder name, rolling back space name", err) // Rollback space name - if rollbackErr := s.spaceDAO.UpdateByID(spaceID, map[string]interface{}{"name": originalName}); rollbackErr != nil { + if rollbackErr := s.spaceDAO.UpdateByID(ctx, dao.DB, spaceID, map[string]interface{}{"name": originalName}); rollbackErr != nil { common.Error("Failed to rollback space name after folder rename failure", rollbackErr) } return nil, common.CodeOperatingError, fmt.Errorf("failed to update folder name: %w", err) @@ -370,21 +370,21 @@ func (s *SkillSpaceService) UpdateSpace(ctx context.Context, spaceID string, ten } if len(updates) > 0 { - if err := s.spaceDAO.UpdateByID(spaceID, updates); err != nil { + if err = s.spaceDAO.UpdateByID(ctx, dao.DB, spaceID, updates); err != nil { return nil, common.CodeOperatingError, fmt.Errorf("failed to update space: %w", err) } } // Refresh space data - space, _ = s.spaceDAO.GetByID(spaceID) + space, _ = s.spaceDAO.GetByID(ctx, dao.DB, spaceID) return space.ToMap(), common.CodeSuccess, nil } // DeleteSpace starts asynchronous deletion of a skills space and returns immediately. // The space status is set to "deleting" and the actual cleanup runs in a background goroutine. -func (s *SkillSpaceService) DeleteSpace(spaceID, tenantID string, docEngine engine.DocEngine, ctx context.Context) (common.ErrorCode, error) { +func (s *SkillSpaceService) DeleteSpace(ctx context.Context, spaceID, tenantID string, docEngine engine.DocEngine) (common.ErrorCode, error) { // Get space regardless of status (could be retrying a failed delete) - space, err := s.spaceDAO.GetByIDAnyStatus(spaceID) + space, err := s.spaceDAO.GetByIDAnyStatus(ctx, dao.DB, spaceID) if err != nil { return common.CodeDataError, fmt.Errorf("space not found") } @@ -407,7 +407,7 @@ func (s *SkillSpaceService) DeleteSpace(spaceID, tenantID string, docEngine engi } // CAS: status must be "1" (active) → "2" (deleting) to prevent concurrent deletes - swapped, err := s.spaceDAO.CASStatus(spaceID, entity.SpaceStatusActive, entity.SpaceStatusDeleting) + swapped, err := s.spaceDAO.CASStatus(ctx, dao.DB, spaceID, entity.SpaceStatusActive, entity.SpaceStatusDeleting) if err != nil { return common.CodeOperatingError, fmt.Errorf("failed to update space status: %w", err) } @@ -419,18 +419,21 @@ func (s *SkillSpaceService) DeleteSpace(spaceID, tenantID string, docEngine engi common.Info("Space marked as deleting, starting async cleanup", zap.String("spaceID", spaceID), zap.String("tenantID", tenantID)) // Launch async deletion in background goroutine - go s.asyncDeleteSpace(spaceID, space.FolderID, tenantID, docEngine, ctx) + go s.asyncDeleteSpace(ctx, spaceID, space.FolderID, tenantID, docEngine) return common.CodeSuccess, nil } // asyncDeleteSpace performs the actual deletion work in the background. // It deletes the search index, removes files via Go FileService, and soft-deletes the space record. -func (s *SkillSpaceService) asyncDeleteSpace(spaceID, folderID, tenantID string, docEngine engine.DocEngine, ctx context.Context) { +func (s *SkillSpaceService) asyncDeleteSpace(ctx context.Context, spaceID, folderID, tenantID string, docEngine engine.DocEngine) { + bgCtx, bgCancel := context.WithTimeout(context.Background(), 120*time.Second) + defer bgCancel() + defer func() { if r := recover(); r != nil { common.Warn("Panic in asyncDeleteSpace, marking space as deleted", zap.Any("recover", r), zap.String("spaceID", spaceID)) - _, _ = s.spaceDAO.CASStatus(spaceID, entity.SpaceStatusDeleting, entity.SpaceStatusDeleted) + _, _ = s.spaceDAO.CASStatus(bgCtx, dao.DB, spaceID, entity.SpaceStatusDeleting, entity.SpaceStatusDeleted) } }() @@ -473,12 +476,12 @@ func (s *SkillSpaceService) asyncDeleteSpace(spaceID, folderID, tenantID string, // Step 3: Soft delete the space record (status "2" → "0") // First, permanently remove any previously deleted spaces with the same tenant+name // to avoid UNIQUE INDEX constraint violation when changing status from "2" to "0" - space, err := s.spaceDAO.GetByIDAnyStatus(spaceID) + space, err := s.spaceDAO.GetByIDAnyStatus(bgCtx, dao.DB, spaceID) if err == nil && space != nil { - _ = s.spaceDAO.DeletePermanentByName(space.TenantID, space.Name) + _ = s.spaceDAO.DeletePermanentByName(bgCtx, dao.DB, space.TenantID, space.Name) } - swapped, err := s.spaceDAO.CASStatus(spaceID, entity.SpaceStatusDeleting, entity.SpaceStatusDeleted) + swapped, err := s.spaceDAO.CASStatus(bgCtx, dao.DB, spaceID, entity.SpaceStatusDeleting, entity.SpaceStatusDeleted) if err != nil { common.Error(fmt.Sprintf("Failed to update space status to deleted, spaceID=%s", spaceID), err) return @@ -536,8 +539,8 @@ func (s *SkillSpaceService) deleteFolderRecursive(ctx context.Context, folderID } // GetSpaceByFolderID retrieves a skills space by its folder ID -func (s *SkillSpaceService) GetSpaceByFolderID(folderID, tenantID string) (map[string]interface{}, common.ErrorCode, error) { - space, err := s.spaceDAO.GetByFolderID(folderID) +func (s *SkillSpaceService) GetSpaceByFolderID(ctx context.Context, folderID, tenantID string) (map[string]interface{}, common.ErrorCode, error) { + space, err := s.spaceDAO.GetByFolderID(ctx, dao.DB, folderID) if err != nil { return nil, common.CodeDataError, fmt.Errorf("space not found for folder") } diff --git a/internal/service/user.go b/internal/service/user.go index 75d7c6cf33..c4541e0891 100644 --- a/internal/service/user.go +++ b/internal/service/user.go @@ -914,11 +914,6 @@ type UserTenantService struct { /** * Returns: * - *UserTenantService: a new UserTenantService instance - * - * Example: - * - * service := NewUserTenantService() - * relations, err := service.GetUserTenantRelationByUserID("user123") */ func NewUserTenantService() *UserTenantService { return &UserTenantService{ @@ -935,37 +930,6 @@ type UserTenantRelation struct { Role string `json:"role"` } -// GetUserTenantRelationByUserID retrieves all user-tenant relationships for a given user ID -/** - * This method returns a list of user-tenant relationships with selected fields: - * - id: the relationship ID - * - user_id: the user ID - * - tenant_id: the tenant ID - * - role: the user's role in the tenant - * - * Parameters: - * - userID: the unique identifier of the user - * - * Returns: - * - []*UserTenantRelation: list of user-tenant relationships - * - error: error if the operation fails, nil otherwise - * - * Example: - * - * service := NewUserTenantService() - * relations, err := service.GetUserTenantRelationByUserID("user123") - * if err != nil { - * log.Printf("Failed to get user tenant relations: %v", err) - * return - * } - * for _, rel := range relations { - * fmt.Printf("User %s has role %s in tenant %s\n", rel.UserID, rel.Role, rel.TenantID) - * } - */ -func (s *UserTenantService) GetUserTenantRelationByUserID(userID string) ([]*UserTenantRelation, error) { - return s.GetUserTenantRelationByUserIDWithContext(context.Background(), userID) -} - // GetUserTenantRelationByUserIDWithContext retrieves all user-tenant relationships for a given user ID with context. func (s *UserTenantService) GetUserTenantRelationByUserIDWithContext(ctx context.Context, userID string) ([]*UserTenantRelation, error) { relations, err := s.userTenantDAO.GetByUserIDWithContext(ctx, userID)