diff --git a/internal/agent/component/agent.go b/internal/agent/component/agent.go index 37b9e7ec3f..4831570796 100644 --- a/internal/agent/component/agent.go +++ b/internal/agent/component/agent.go @@ -753,7 +753,7 @@ func (c *AgentComponent) invokeNow(ctx context.Context, db *gorm.DB, inputs map[ } var err error - p.ModelID, p.Driver, p.APIKey, p.BaseURL, err = resolveChatModelRef(ctx, p.ModelID, p.Driver, p.APIKey, p.BaseURL) + p.ModelID, p.Driver, p.APIKey, p.BaseURL, err = resolveChatModelRef(ctx, db, p.ModelID, p.Driver, p.APIKey, p.BaseURL) if err != nil { return nil, err } diff --git a/internal/agent/component/browser.go b/internal/agent/component/browser.go index 21e730a8e2..4b524ad396 100644 --- a/internal/agent/component/browser.go +++ b/internal/agent/component/browser.go @@ -240,7 +240,7 @@ func (p *browserParam) AsDict() map[string]any { } // BrowserComponent is the canvas Browser node. Owns its static -// param; delegates the multi-step agent run to StagehandInvoker. +// param; delegates the multistep agent run to StagehandInvoker. type BrowserComponent struct { name string param browserParam @@ -250,10 +250,10 @@ type BrowserComponent struct { func NewBrowserComponent(params map[string]any) (Component, error) { p := &browserParam{} if err := p.Update(params); err != nil { - return nil, fmt.Errorf("Browser: param update: %w", err) + return nil, fmt.Errorf("browser: param update: %w", err) } if err := p.Check(); err != nil { - return nil, fmt.Errorf("Browser: param check: %w", err) + return nil, fmt.Errorf("browser: param check: %w", err) } return &BrowserComponent{ name: componentNameBrowser, @@ -285,15 +285,15 @@ func (b *BrowserComponent) Name() string { return b.name } func (b *BrowserComponent) Invoke(ctx context.Context, db *gorm.DB, inputs map[string]any) (map[string]any, error) { state, _, err := runtime.GetStateFromContext[*runtime.CanvasState](ctx) if err != nil { - return nil, fmt.Errorf("Browser: %w", err) + return nil, fmt.Errorf("browser: %w", err) } if state == nil { - return nil, errors.New("Browser: nil canvas state") + return nil, errors.New("browser: nil canvas state") } tenantID, _ := state.Sys["tenant_id"].(string) if tenantID == "" { - return nil, errors.New("Browser: tenant_id missing from canvas state (state.Sys[\"tenant_id\"])") + return nil, errors.New("browser: tenant_id missing from canvas state (state.Sys[\"tenant_id\"])") } // 1. Resolve prompts template. @@ -303,13 +303,13 @@ func (b *BrowserComponent) Invoke(ctx context.Context, db *gorm.DB, inputs map[s } resolvedPrompts, err := runtime.ResolveTemplate(prompts, state) if err != nil { - return nil, fmt.Errorf("Browser: resolve prompts template: %w", err) + return nil, fmt.Errorf("browser: resolve prompts template: %w", err) } // 2. Look up tenant model config. - providerName, modelName, apiKey, baseURL, err := resolveBrowserLLM(tenantID, b.param.LLMID) + providerName, modelName, apiKey, baseURL, err := resolveBrowserLLM(ctx, db, tenantID, b.param.LLMID) if err != nil { - return nil, fmt.Errorf("Browser: tenant llm lookup (%q): %w", b.param.LLMID, err) + return nil, fmt.Errorf("browser: tenant llm lookup (%q): %w", b.param.LLMID, err) } baseURL = strings.TrimSpace(baseURL) @@ -329,14 +329,14 @@ func (b *BrowserComponent) Invoke(ctx context.Context, db *gorm.DB, inputs map[s invoker := getDefaultStagehandInvoker() rawJSON, err := invoker.RunExtract(ctx, req) if err != nil { - return nil, fmt.Errorf("Browser: stagehand extract (model=%q, base_url=%s): %w", + return nil, fmt.Errorf("browser: stagehand extract (model=%q, base_url=%s): %w", req.ModelName, browserBaseURLForLog(req.BaseURL), err) } // 5. Unmarshal the JSON-string result to get the plain text. var content string - if err := json.Unmarshal([]byte(rawJSON), &content); err != nil { - return nil, fmt.Errorf("Browser: unmarshal extract result: %w", err) + if err = json.Unmarshal([]byte(rawJSON), &content); err != nil { + return nil, fmt.Errorf("browser: unmarshal extract result: %w", err) } // 6. Build the output map. @@ -444,7 +444,7 @@ func (b *BrowserComponent) Outputs() map[string]string { // Tests override the lookup via `tenantLLMLookupForTest` (a // package-level function variable) so they don't need a real DB. // Production code leaves the variable unset. -func resolveBrowserLLM(tenantID, llmID string) (providerName, modelName, apiKey, baseURL string, err error) { +func resolveBrowserLLM(ctx context.Context, db *gorm.DB, tenantID, llmID string) (providerName, modelName, apiKey, baseURL string, err error) { if tenantLLMLookupForTest != nil { oldModelName, factory := resolveLLMID(llmID) apiKey, baseURL, err = tenantLLMLookupForTest(tenantID, oldModelName, factory) @@ -452,7 +452,7 @@ func resolveBrowserLLM(tenantID, llmID string) (providerName, modelName, apiKey, return factory, oldModelName, apiKey, baseURL, err } - providerName, modelName, apiKey, baseURL, err = resolveTenantModelBrowserLLM(tenantID, llmID) + providerName, modelName, apiKey, baseURL, err = resolveTenantModelBrowserLLM(ctx, db, tenantID, llmID) if err == nil { baseURL = browserOpenAICompatibleBaseURL(baseURL, providerName) return providerName, modelName, apiKey, baseURL, nil @@ -460,7 +460,7 @@ func resolveBrowserLLM(tenantID, llmID string) (providerName, modelName, apiKey, modelErr := err oldModelName, factory := resolveLLMID(llmID) - apiKey, baseURL, oldErr := resolveTenantLLM(tenantID, oldModelName, factory) + apiKey, baseURL, oldErr := resolveTenantLLM(ctx, db, tenantID, oldModelName, factory) if oldErr == nil { baseURL = browserOpenAICompatibleBaseURL(baseURL, factory) return factory, oldModelName, apiKey, baseURL, nil @@ -468,8 +468,8 @@ func resolveBrowserLLM(tenantID, llmID string) (providerName, modelName, apiKey, return "", "", "", "", fmt.Errorf("tenant_model lookup: %v; tenant_llm fallback: %w", modelErr, oldErr) } -func resolveTenantModelBrowserLLM(tenantID, modelID string) (providerName, modelName, apiKey, baseURL string, err error) { - modelRow, err := dao.NewTenantModelDAO().GetByID(modelID) +func resolveTenantModelBrowserLLM(ctx context.Context, db *gorm.DB, tenantID, modelID string) (providerName, modelName, apiKey, baseURL string, err error) { + modelRow, err := dao.NewTenantModelDAO().GetByID(ctx, db, modelID) if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return "", "", "", "", err @@ -483,7 +483,7 @@ func resolveTenantModelBrowserLLM(tenantID, modelID string) (providerName, model return "", "", "", "", fmt.Errorf("tenant model id=%s cannot be used as %s model", modelID, entity.ModelTypeChat.String()) } - provider, err := dao.NewTenantModelProviderDAO().GetByID(modelRow.ProviderID) + provider, err := dao.NewTenantModelProviderDAO().GetByID(ctx, db, modelRow.ProviderID) if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return "", "", "", "", fmt.Errorf("provider id=%s not found for model id=%s", modelRow.ProviderID, modelID) @@ -497,7 +497,7 @@ func resolveTenantModelBrowserLLM(tenantID, modelID string) (providerName, model return "", "", "", "", fmt.Errorf("tenant %s has no access to provider owned by tenant %s", tenantID, provider.TenantID) } - instance, err := dao.NewTenantModelInstanceDAO().GetByID(modelRow.InstanceID) + instance, err := dao.NewTenantModelInstanceDAO().GetByID(ctx, db, modelRow.InstanceID) if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return "", "", "", "", fmt.Errorf("instance id=%s not found for model id=%s", modelRow.InstanceID, modelID) @@ -511,7 +511,7 @@ func resolveTenantModelBrowserLLM(tenantID, modelID string) (providerName, model apiKey = instance.APIKey if strings.TrimSpace(instance.Extra) != "" { var extra map[string]string - if err := json.Unmarshal([]byte(instance.Extra), &extra); err != nil { + if err = json.Unmarshal([]byte(instance.Extra), &extra); err != nil { return "", "", "", "", err } baseURL = extra["base_url"] @@ -537,18 +537,18 @@ func browserOpenAICompatibleBaseURL(baseURL, provider string) string { // // TODO(v2): this helper can move to `internal/dao` so the LLM // component (`llm.go`) and other future components can share it. -func resolveTenantLLM(tenantID, modelName, factory string) (apiKey, baseURL string, err error) { - dao := dao.NewTenantLLMDAO() +func resolveTenantLLM(ctx context.Context, db *gorm.DB, tenantID, modelName, factory string) (apiKey, baseURL string, err error) { + llmDAO := dao.NewTenantLLMDAO() var ( row *entity.TenantLLM ) if factory != "" { - row, err = dao.GetByTenantFactoryAndModelName(tenantID, factory, modelName) + row, err = llmDAO.GetByTenantFactoryAndModelName(ctx, db, tenantID, factory, modelName) } else { // No factory suffix on llm_id; fall back to a single-key // lookup (errors if the model is registered under multiple // factories — caller must use the explicit form). - row, err = dao.GetByTenantAndModelName(tenantID, "", modelName) + row, err = llmDAO.GetByTenantAndModelName(ctx, db, tenantID, "", modelName) } if err != nil { return "", "", err diff --git a/internal/agent/component/browser_test.go b/internal/agent/component/browser_test.go index 17f685cfed..e51f96373a 100644 --- a/internal/agent/component/browser_test.go +++ b/internal/agent/component/browser_test.go @@ -249,7 +249,8 @@ func TestResolveBrowserLLM_ResolvesTenantModelID(t *testing.T) { tenantLLMLookupForTest = nil t.Cleanup(func() { tenantLLMLookupForTest = prevLookup }) - provider, model, apiKey, baseURL, err := resolveBrowserLLM("tenant-1", "tenant-model-1") + ctx := t.Context() + provider, model, apiKey, baseURL, err := resolveBrowserLLM(ctx, db, "tenant-1", "tenant-model-1") if err != nil { t.Fatalf("resolveBrowserLLM: %v", err) } diff --git a/internal/agent/component/categorize.go b/internal/agent/component/categorize.go index e482c7a15f..8b9367235e 100644 --- a/internal/agent/component/categorize.go +++ b/internal/agent/component/categorize.go @@ -68,7 +68,7 @@ func (c *CategorizeComponent) Name() string { return "Categorize" } func (c *CategorizeComponent) Invoke(ctx context.Context, db *gorm.DB, inputs map[string]any) (map[string]any, error) { p := mergeCategorizeParam(c.param, inputs) var err error - p.ModelID, p.Driver, p.APIKey, p.BaseURL, err = resolveChatModelRef(ctx, p.ModelID, p.Driver, p.APIKey, p.BaseURL) + p.ModelID, p.Driver, p.APIKey, p.BaseURL, err = resolveChatModelRef(ctx, db, p.ModelID, p.Driver, p.APIKey, p.BaseURL) if err != nil { return nil, err } diff --git a/internal/agent/component/llm.go b/internal/agent/component/llm.go index 5101221ed9..aaca4985d1 100644 --- a/internal/agent/component/llm.go +++ b/internal/agent/component/llm.go @@ -321,7 +321,7 @@ func (c *LLMComponent) Invoke(ctx context.Context, db *gorm.DB, inputs map[strin // reference selected in the agent canvas is passed verbatim to the LLM // driver, causing 400s for custom-added models. var err error - p.ModelID, p.Driver, p.APIKey, p.BaseURL, err = resolveChatModelRef(ctx, p.ModelID, p.Driver, p.APIKey, p.BaseURL) + p.ModelID, p.Driver, p.APIKey, p.BaseURL, err = resolveChatModelRef(ctx, db, p.ModelID, p.Driver, p.APIKey, p.BaseURL) if err != nil { return nil, fmt.Errorf("component: LLM.Invoke: resolve model: %w", err) } diff --git a/internal/agent/component/llm_credentials.go b/internal/agent/component/llm_credentials.go index 71f26b7a46..f675ac6529 100644 --- a/internal/agent/component/llm_credentials.go +++ b/internal/agent/component/llm_credentials.go @@ -20,7 +20,7 @@ import ( // driver/model pair when the canvas DSL omitted them. It first checks the old // tenant_llm table, then falls back to tenant_model_provider + // tenant_model_instance when the composite llm_id carries an instance name. -func resolveTenantLLMConfig(ctx context.Context, driver, modelID, apiKey, baseURL, originalModelID string) (string, string) { +func resolveTenantLLMConfig(ctx context.Context, db *gorm.DB, driver, modelID, apiKey, baseURL, originalModelID string) (string, string) { if apiKey != "" || driver == "" || modelID == "" { return apiKey, baseURL } @@ -35,19 +35,19 @@ func resolveTenantLLMConfig(ctx context.Context, driver, modelID, apiKey, baseUR return apiKey, baseURL } - if resolvedKey, resolvedBaseURL, ok := resolveTenantLLMCredentials(tid, driver, modelID, baseURL); ok { + if resolvedKey, resolvedBaseURL, ok := resolveTenantLLMCredentials(ctx, db, tid, driver, modelID, baseURL); ok { return resolvedKey, resolvedBaseURL } if originalModelID == "" { return apiKey, baseURL } - if resolvedKey, resolvedBaseURL, ok := resolveTenantModelInstanceCredentials(tid, originalModelID, baseURL); ok { + if resolvedKey, resolvedBaseURL, ok := resolveTenantModelInstanceCredentials(ctx, db, tid, originalModelID, baseURL); ok { return resolvedKey, resolvedBaseURL } return apiKey, baseURL } -func resolveChatModelRef(ctx context.Context, modelID, driver, apiKey, baseURL string) (string, string, string, string, error) { +func resolveChatModelRef(ctx context.Context, db *gorm.DB, modelID, driver, apiKey, baseURL string) (string, string, string, string, error) { originalModelID := modelID if driver == "" && modelID != "" { if m, prov, ok := agentProviderLastSegmentSplit(modelID); ok { @@ -56,7 +56,7 @@ func resolveChatModelRef(ctx context.Context, modelID, driver, apiKey, baseURL s } } if driver == "" && modelID != "" { - resolvedModelID, resolvedDriver, resolvedAPIKey, resolvedBaseURL, ok, err := resolveTenantChatModelByID(ctx, modelID, apiKey, baseURL) + resolvedModelID, resolvedDriver, resolvedAPIKey, resolvedBaseURL, ok, err := resolveTenantChatModelByID(ctx, db, modelID, apiKey, baseURL) if err != nil { return "", "", "", "", err } @@ -67,15 +67,15 @@ func resolveChatModelRef(ctx context.Context, modelID, driver, apiKey, baseURL s baseURL = resolvedBaseURL } } - apiKey, baseURL = resolveTenantLLMConfig(ctx, driver, modelID, apiKey, baseURL, originalModelID) + apiKey, baseURL = resolveTenantLLMConfig(ctx, db, driver, modelID, apiKey, baseURL, originalModelID) return modelID, driver, apiKey, baseURL, nil } // resolveTenantLLMCredentials looks up the old tenant_llm table for the given // tenant / factory / model. Returns true when credentials were found. -func resolveTenantLLMCredentials(tid, driver, modelID, baseURL string) (string, string, bool) { +func resolveTenantLLMCredentials(ctx context.Context, db *gorm.DB, tid, driver, modelID, baseURL string) (string, string, bool) { common.Debug("llm credentials: tenant_llm lookup", zap.String("tid", tid), zap.String("factory", driver), zap.String("model", modelID)) - row, err := dao.NewTenantLLMDAO().GetByTenantFactoryAndModelName(tid, driver, modelID) + row, err := dao.NewTenantLLMDAO().GetByTenantFactoryAndModelName(ctx, db, tid, driver, modelID) if err != nil { common.Debug("llm credentials: tenant_llm lookup", zap.Error(err)) return "", baseURL, false @@ -101,7 +101,7 @@ func resolveTenantLLMCredentials(tid, driver, modelID, baseURL string) (string, // resolveTenantModelInstanceCredentials attempts to resolve llm credentials // through tenant_model_provider + tenant_model_instance using the original // composite llm_id (which still carries the instance name). -func resolveTenantModelInstanceCredentials(tid, compositeLLMID, baseURL string) (string, string, bool) { +func resolveTenantModelInstanceCredentials(ctx context.Context, db *gorm.DB, tid, compositeLLMID, baseURL string) (string, string, bool) { modelName, instanceName, providerName := parseLLMIDParts(compositeLLMID) if instanceName == "" { common.Debug("llm credentials: new-table fallback skipped: no instance name", zap.String("composite_llm_id", compositeLLMID)) @@ -114,16 +114,16 @@ func resolveTenantModelInstanceCredentials(tid, compositeLLMID, baseURL string) zap.String("model", modelName), zap.String("instance", instanceName)) - provider, err := dao.NewTenantModelProviderDAO().GetByTenantIDAndProviderName(tid, providerName) + provider, err := dao.NewTenantModelProviderDAO().GetByTenantIDAndProviderName(ctx, db, tid, providerName) if err != nil || provider == nil { common.Debug("llm credentials: new-table fallback: provider not found", zap.String("provider", providerName), zap.Error(err)) return "", baseURL, false } - instance, err := dao.NewTenantModelInstanceDAO().GetByProviderIDAndInstanceName(provider.ID, instanceName) + instance, err := dao.NewTenantModelInstanceDAO().GetByProviderIDAndInstanceName(ctx, db, provider.ID, instanceName) if err != nil || instance == nil { if instanceName == "default" { - if fallback := findSoleActiveProviderInstance(provider.ID); fallback != nil { + if fallback := findSoleActiveProviderInstance(ctx, db, provider.ID); fallback != nil { common.Debug("llm credentials: new-table fallback: remapped default instance to sole active instance", zap.String("instance", fallback.InstanceName), zap.String("provider", providerName)) @@ -158,8 +158,8 @@ func resolveTenantModelInstanceCredentials(tid, compositeLLMID, baseURL string) return apiKey, baseURL, apiKey != "" } -func findSoleActiveProviderInstance(providerID string) *entity.TenantModelInstance { - instances, err := dao.NewTenantModelInstanceDAO().GetAllInstancesByProviderID(providerID) +func findSoleActiveProviderInstance(ctx context.Context, db *gorm.DB, providerID string) *entity.TenantModelInstance { + instances, err := dao.NewTenantModelInstanceDAO().GetAllInstancesByProviderID(ctx, db, providerID) if err != nil { common.Debug("llm credentials: list provider instances", zap.Error(err)) return nil @@ -180,7 +180,7 @@ func findSoleActiveProviderInstance(providerID string) *entity.TenantModelInstan return active[0] } -func resolveTenantChatModelByID(ctx context.Context, modelRef, apiKey, baseURL string) (string, string, string, string, bool, error) { +func resolveTenantChatModelByID(ctx context.Context, db *gorm.DB, modelRef, apiKey, baseURL string) (string, string, string, string, bool, error) { if !isBareTenantModelID(modelRef) { return "", "", apiKey, baseURL, false, nil } @@ -193,7 +193,7 @@ func resolveTenantChatModelByID(ctx context.Context, modelRef, apiKey, baseURL s return "", "", apiKey, baseURL, false, nil } - modelName, provider, modelKey, modelBaseURL, ok, err := resolveTenantChatModelByTenantModelID(tid, modelRef, apiKey, baseURL) + modelName, provider, modelKey, modelBaseURL, ok, err := resolveTenantChatModelByTenantModelID(ctx, db, tid, modelRef, apiKey, baseURL) if err != nil { return "", "", apiKey, baseURL, false, err } @@ -201,7 +201,7 @@ func resolveTenantChatModelByID(ctx context.Context, modelRef, apiKey, baseURL s return modelName, provider, modelKey, modelBaseURL, true, nil } - modelName, provider, instanceKey, instanceBaseURL, ok, err := resolveTenantChatModelByInstanceID(tid, modelRef, apiKey, baseURL) + modelName, provider, instanceKey, instanceBaseURL, ok, err := resolveTenantChatModelByInstanceID(ctx, db, tid, modelRef, apiKey, baseURL) if err != nil { return "", "", apiKey, baseURL, false, err } @@ -211,8 +211,8 @@ func resolveTenantChatModelByID(ctx context.Context, modelRef, apiKey, baseURL s return "", "", apiKey, baseURL, false, fmt.Errorf("tenant chat model id %q not found", modelRef) } -func resolveTenantChatModelByTenantModelID(tid, modelID, apiKey, baseURL string) (string, string, string, string, bool, error) { - model, err := dao.NewTenantModelDAO().GetByID(modelID) +func resolveTenantChatModelByTenantModelID(ctx context.Context, db *gorm.DB, tid, modelID, apiKey, baseURL string) (string, string, string, string, bool, error) { + model, err := dao.NewTenantModelDAO().GetByID(ctx, db, modelID) if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return "", "", apiKey, baseURL, false, nil @@ -226,7 +226,7 @@ func resolveTenantChatModelByTenantModelID(tid, modelID, apiKey, baseURL string) return "", "", apiKey, baseURL, false, fmt.Errorf("tenant model id %s cannot be used as chat model", modelID) } - provider, err := dao.NewTenantModelProviderDAO().GetByID(model.ProviderID) + provider, err := dao.NewTenantModelProviderDAO().GetByID(ctx, db, model.ProviderID) if err != nil { return "", "", apiKey, baseURL, false, fmt.Errorf("resolve provider for tenant model id %q: %w", modelID, err) } @@ -234,7 +234,7 @@ func resolveTenantChatModelByTenantModelID(tid, modelID, apiKey, baseURL string) return "", "", apiKey, baseURL, false, fmt.Errorf("tenant %s has no access to model id %s", tid, modelID) } - instance, err := dao.NewTenantModelInstanceDAO().GetByID(model.InstanceID) + instance, err := dao.NewTenantModelInstanceDAO().GetByID(ctx, db, model.InstanceID) if err != nil { return "", "", apiKey, baseURL, false, fmt.Errorf("resolve instance for tenant model id %q: %w", modelID, err) } @@ -242,15 +242,15 @@ func resolveTenantChatModelByTenantModelID(tid, modelID, apiKey, baseURL string) return model.ModelName, provider.ProviderName, apiKey, baseURL, true, nil } -func resolveTenantChatModelByInstanceID(tid, instanceID, apiKey, baseURL string) (string, string, string, string, bool, error) { - instance, err := dao.NewTenantModelInstanceDAO().GetByID(instanceID) +func resolveTenantChatModelByInstanceID(ctx context.Context, db *gorm.DB, tid, instanceID, apiKey, baseURL string) (string, string, string, string, bool, error) { + instance, err := dao.NewTenantModelInstanceDAO().GetByID(ctx, db, instanceID) if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return "", "", apiKey, baseURL, false, nil } return "", "", apiKey, baseURL, false, fmt.Errorf("resolve tenant model instance id %q: %w", instanceID, err) } - provider, err := dao.NewTenantModelProviderDAO().GetByID(instance.ProviderID) + provider, err := dao.NewTenantModelProviderDAO().GetByID(ctx, db, instance.ProviderID) if err != nil { return "", "", apiKey, baseURL, false, fmt.Errorf("resolve provider for tenant model instance id %q: %w", instanceID, err) } @@ -258,7 +258,7 @@ func resolveTenantChatModelByInstanceID(tid, instanceID, apiKey, baseURL string) return "", "", apiKey, baseURL, false, fmt.Errorf("tenant %s has no access to model instance id %s", tid, instanceID) } - models, err := dao.NewTenantModelDAO().GetModelsByInstanceID(instance.ID) + models, err := dao.NewTenantModelDAO().GetModelsByInstanceID(ctx, db, instance.ID) if err != nil { return "", "", apiKey, baseURL, false, fmt.Errorf("resolve models for tenant model instance id %q: %w", instanceID, err) } diff --git a/internal/dao/tenant_llm.go b/internal/dao/tenant_llm.go index 4325039c91..aaeba1dc4a 100644 --- a/internal/dao/tenant_llm.go +++ b/internal/dao/tenant_llm.go @@ -17,8 +17,11 @@ package dao import ( + "context" "fmt" "ragflow/internal/entity" + + "gorm.io/gorm" ) // TenantLLMDAO tenant LLM data access object @@ -30,9 +33,9 @@ func NewTenantLLMDAO() *TenantLLMDAO { } // GetByID get tenant LLM by primary key ID -func (dao *TenantLLMDAO) GetByID(id int64) (*entity.TenantLLM, error) { +func (dao *TenantLLMDAO) GetByID(ctx context.Context, db *gorm.DB, id int64) (*entity.TenantLLM, error) { var tenantLLM entity.TenantLLM - err := DB.Where("id = ?", id).First(&tenantLLM).Error + err := db.WithContext(ctx).Where("id = ?", id).First(&tenantLLM).Error if err != nil { return nil, err } @@ -40,9 +43,9 @@ func (dao *TenantLLMDAO) GetByID(id int64) (*entity.TenantLLM, error) { } // GetByTenantAndModelName get tenant LLM by tenant ID and model name -func (dao *TenantLLMDAO) GetByTenantAndModelName(tenantID, providerName string, modelName string) (*entity.TenantLLM, error) { +func (dao *TenantLLMDAO) GetByTenantAndModelName(ctx context.Context, db *gorm.DB, tenantID, providerName string, modelName string) (*entity.TenantLLM, error) { var tenantLLM entity.TenantLLM - err := DB.Where("tenant_id = ? AND llm_factory = ? AND llm_name = ?", tenantID, providerName, modelName).First(&tenantLLM).Error + err := db.WithContext(ctx).Where("tenant_id = ? AND llm_factory = ? AND llm_name = ?", tenantID, providerName, modelName).First(&tenantLLM).Error if err != nil { return nil, err } @@ -50,10 +53,10 @@ func (dao *TenantLLMDAO) GetByTenantAndModelName(tenantID, providerName string, } // GetByTenantNameAndType get tenant LLM by tenant ID, model name, and model type -func (dao *TenantLLMDAO) GetByTenantNameAndType(tenantID, modelName string, modelType entity.ModelType) (*entity.TenantLLM, error) { +func (dao *TenantLLMDAO) GetByTenantNameAndType(ctx context.Context, db *gorm.DB, tenantID, modelName string, modelType entity.ModelType) (*entity.TenantLLM, error) { var tenantLLM entity.TenantLLM // tenant_llm.model_type is a VARCHAR column, so convert to string. - err := DB.Where("tenant_id = ? AND llm_name = ? AND model_type = ?", tenantID, modelName, modelType.String()).First(&tenantLLM).Error + err := db.WithContext(ctx).Where("tenant_id = ? AND llm_name = ? AND model_type = ?", tenantID, modelName, modelType.String()).First(&tenantLLM).Error if err != nil { return nil, err } @@ -61,10 +64,10 @@ func (dao *TenantLLMDAO) GetByTenantNameAndType(tenantID, modelName string, mode } // GetByTenantAndType get tenant LLM by tenant ID and model type -func (dao *TenantLLMDAO) GetByTenantAndType(tenantID string, modelType entity.ModelType) (*entity.TenantLLM, error) { +func (dao *TenantLLMDAO) GetByTenantAndType(ctx context.Context, db *gorm.DB, tenantID string, modelType entity.ModelType) (*entity.TenantLLM, error) { var tenantLLM entity.TenantLLM // tenant_llm.model_type is a VARCHAR column, so convert to string. - err := DB.Where("tenant_id = ? AND model_type = ?", tenantID, modelType.String()).First(&tenantLLM).Error + err := db.WithContext(ctx).Where("tenant_id = ? AND model_type = ?", tenantID, modelType.String()).First(&tenantLLM).Error if err != nil { return nil, err } @@ -72,10 +75,10 @@ func (dao *TenantLLMDAO) GetByTenantAndType(tenantID string, modelType entity.Mo } // GetByTenantAndFactory get tenant LLM by tenant ID, model type and factory -func (dao *TenantLLMDAO) GetByTenantAndFactory(tenantID string, modelType entity.ModelType, factory string) (*entity.TenantLLM, error) { +func (dao *TenantLLMDAO) GetByTenantAndFactory(ctx context.Context, db *gorm.DB, tenantID string, modelType entity.ModelType, factory string) (*entity.TenantLLM, error) { var tenantLLM entity.TenantLLM // tenant_llm.model_type is a VARCHAR column, so convert to string. - err := DB.Where("tenant_id = ? AND model_type = ? AND llm_factory = ?", tenantID, modelType.String(), factory).First(&tenantLLM).Error + err := db.WithContext(ctx).Where("tenant_id = ? AND model_type = ? AND llm_factory = ?", tenantID, modelType.String(), factory).First(&tenantLLM).Error if err != nil { return nil, err } @@ -83,9 +86,9 @@ func (dao *TenantLLMDAO) GetByTenantAndFactory(tenantID string, modelType entity } // ListByTenant list all tenant LLMs for a tenant -func (dao *TenantLLMDAO) ListByTenant(tenantID string) ([]entity.TenantLLM, error) { +func (dao *TenantLLMDAO) ListByTenant(ctx context.Context, db *gorm.DB, tenantID string) ([]entity.TenantLLM, error) { var tenantLLMs []entity.TenantLLM - err := DB.Where("tenant_id = ?", tenantID).Find(&tenantLLMs).Error + err := db.WithContext(ctx).Where("tenant_id = ?", tenantID).Find(&tenantLLMs).Error if err != nil { return nil, err } @@ -93,35 +96,35 @@ func (dao *TenantLLMDAO) ListByTenant(tenantID string) ([]entity.TenantLLM, erro } // GetByTenantFactoryAndModelName get tenant LLM by tenant ID, factory and model name -func (dao *TenantLLMDAO) GetByTenantFactoryAndModelName(tenantID, factory, modelName string) (*entity.TenantLLM, error) { +func (dao *TenantLLMDAO) GetByTenantFactoryAndModelName(ctx context.Context, db *gorm.DB, tenantID, factory, modelName string) (*entity.TenantLLM, error) { var tenantLLM entity.TenantLLM - err := DB.Where("tenant_id = ? AND llm_factory = ? AND llm_name = ?", tenantID, factory, modelName).First(&tenantLLM).Error + err := db.WithContext(ctx).Where("tenant_id = ? AND llm_factory = ? AND llm_name = ?", tenantID, factory, modelName).First(&tenantLLM).Error if err != nil { return nil, err } return &tenantLLM, nil } -// Create create a new tenant LLM record -func (dao *TenantLLMDAO) Create(tenantLLM *entity.TenantLLM) error { - return DB.Create(tenantLLM).Error +// Create a new tenant LLM record +func (dao *TenantLLMDAO) Create(ctx context.Context, db *gorm.DB, tenantLLM *entity.TenantLLM) error { + return db.WithContext(ctx).Create(tenantLLM).Error } -// Update update an existing tenant LLM record -func (dao *TenantLLMDAO) Update(tenantLLM *entity.TenantLLM) error { - return DB.Save(tenantLLM).Error +// Update an existing tenant LLM record +func (dao *TenantLLMDAO) Update(ctx context.Context, db *gorm.DB, tenantLLM *entity.TenantLLM) error { + return db.WithContext(ctx).Save(tenantLLM).Error } -// Delete delete a tenant LLM record by tenant ID, factory and model name -func (dao *TenantLLMDAO) Delete(tenantID, factory, modelName string) error { - return DB.Where("tenant_id = ? AND llm_factory = ? AND llm_name = ?", tenantID, factory, modelName).Delete(&entity.TenantLLM{}).Error +// Delete a tenant LLM record by tenant ID, factory and model name +func (dao *TenantLLMDAO) Delete(ctx context.Context, db *gorm.DB, tenantID, factory, modelName string) error { + return db.WithContext(ctx).Where("tenant_id = ? AND llm_factory = ? AND llm_name = ?", tenantID, factory, modelName).Delete(&entity.TenantLLM{}).Error } // GetMyLLMs get tenant LLMs with factory details -func (dao *TenantLLMDAO) GetMyLLMs(tenantID string) ([]entity.MyLLM, error) { +func (dao *TenantLLMDAO) GetMyLLMs(ctx context.Context, db *gorm.DB, tenantID string) ([]entity.MyLLM, error) { var myLLMs []entity.MyLLM - err := DB.Table("tenant_llm tl"). + err := db.WithContext(ctx).Table("tenant_llm tl"). Select("tl.id, tl.llm_factory, lf.logo, lf.tags, tl.model_type, tl.llm_name, tl.used_tokens, tl.status"). Joins("JOIN llm_factories lf ON tl.llm_factory = lf.name"). Where("tl.tenant_id = ? AND tl.api_key IS NOT NULL", tenantID). @@ -133,9 +136,9 @@ func (dao *TenantLLMDAO) GetMyLLMs(tenantID string) ([]entity.MyLLM, error) { } // ListValidByTenant lists valid tenant LLMs for a tenant -func (dao *TenantLLMDAO) ListValidByTenant(tenantID string) ([]*entity.TenantLLM, error) { +func (dao *TenantLLMDAO) ListValidByTenant(ctx context.Context, db *gorm.DB, tenantID string) ([]*entity.TenantLLM, error) { var tenantLLMs []*entity.TenantLLM - err := DB.Where("tenant_id = ? AND api_key IS NOT NULL AND api_key != ? AND status = ?", tenantID, "", "1").Find(&tenantLLMs).Error + err := db.WithContext(ctx).Where("tenant_id = ? AND api_key IS NOT NULL AND api_key != ? AND status = ?", tenantID, "", "1").Find(&tenantLLMs).Error if err != nil { return nil, err } @@ -143,9 +146,9 @@ func (dao *TenantLLMDAO) ListValidByTenant(tenantID string) ([]*entity.TenantLLM } // ListAllByTenant lists all tenant LLMs for a tenant -func (dao *TenantLLMDAO) ListAllByTenant(tenantID string) ([]*entity.TenantLLM, error) { +func (dao *TenantLLMDAO) ListAllByTenant(ctx context.Context, db *gorm.DB, tenantID string) ([]*entity.TenantLLM, error) { var tenantLLMs []*entity.TenantLLM - err := DB.Where("tenant_id = ?", tenantID).Find(&tenantLLMs).Error + err := db.WithContext(ctx).Where("tenant_id = ?", tenantID).Find(&tenantLLMs).Error if err != nil { return nil, err } @@ -153,16 +156,16 @@ func (dao *TenantLLMDAO) ListAllByTenant(tenantID string) ([]*entity.TenantLLM, } // InsertMany inserts multiple tenant LLM records -func (dao *TenantLLMDAO) InsertMany(tenantLLMs []*entity.TenantLLM) error { +func (dao *TenantLLMDAO) InsertMany(ctx context.Context, db *gorm.DB, tenantLLMs []*entity.TenantLLM) error { if len(tenantLLMs) == 0 { return nil } - return DB.Create(&tenantLLMs).Error + return db.WithContext(ctx).Create(&tenantLLMs).Error } // DeleteByTenantID deletes all tenant LLM records by tenant ID (hard delete) -func (dao *TenantLLMDAO) DeleteByTenantID(tenantID string) (int64, error) { - result := DB.Unscoped().Where("tenant_id = ?", tenantID).Delete(&entity.TenantLLM{}) +func (dao *TenantLLMDAO) DeleteByTenantID(ctx context.Context, db *gorm.DB, tenantID string) (int64, error) { + result := db.WithContext(ctx).Unscoped().Where("tenant_id = ?", tenantID).Delete(&entity.TenantLLM{}) return result.RowsAffected, result.Error } @@ -183,7 +186,7 @@ func (dao *TenantLLMDAO) DeleteByTenantID(tenantID string) (int64, error) { // // modelName, factory := splitModelNameAndFactory("gpt-4@OpenAI") // // Returns: "gpt-4", "OpenAI" -func splitModelNameAndFactory(modelName string) (string, string) { +func splitModelNameAndFactory(ctx context.Context, db *gorm.DB, modelName string) (string, string) { // Split by "@" separator // Handle cases like "model@factory" or "model@sub@factory" lastAtIndex := -1 @@ -206,7 +209,7 @@ func splitModelNameAndFactory(modelName string) (string, string) { // Validate if factory exists in llm_factories table // This matches Python's logic of checking against model providers var factoryCount int64 - DB.Model(&entity.LLMFactories{}).Where("name = ?", factory).Count(&factoryCount) + db.WithContext(ctx).Model(&entity.LLMFactories{}).Where("name = ?", factory).Count(&factoryCount) // If factory doesn't exist in database, treat the whole string as model name if factoryCount == 0 { @@ -235,21 +238,21 @@ func splitModelNameAndFactory(modelName string) (string, string) { // // // Model name with factory prefix // tenantLLM, err := dao.GetByTenantIDAndLLMName("tenant123", "gpt-4@OpenAI") -func (dao *TenantLLMDAO) GetByTenantIDAndLLMName(tenantID string, llmName string) (*entity.TenantLLM, error) { +func (dao *TenantLLMDAO) GetByTenantIDAndLLMName(ctx context.Context, db *gorm.DB, tenantID string, llmName string) (*entity.TenantLLM, error) { var tenantLLM entity.TenantLLM // Split model name and factory from the combined format - modelName, factory := splitModelNameAndFactory(llmName) + modelName, factory := splitModelNameAndFactory(ctx, db, llmName) // First attempt: try to find with model name only - err := DB.Where("tenant_id = ? AND llm_name = ?", tenantID, modelName).First(&tenantLLM).Error + err := db.WithContext(ctx).Where("tenant_id = ? AND llm_name = ?", tenantID, modelName).First(&tenantLLM).Error if err == nil { return &tenantLLM, nil } // Second attempt: if factory is specified, try with both model name and factory if factory != "" { - err = DB.Where("tenant_id = ? AND llm_name = ? AND llm_factory = ?", tenantID, modelName, factory).First(&tenantLLM).Error + err = db.WithContext(ctx).Where("tenant_id = ? AND llm_name = ? AND llm_factory = ?", tenantID, modelName, factory).First(&tenantLLM).Error if err == nil { return &tenantLLM, nil } @@ -258,7 +261,7 @@ func (dao *TenantLLMDAO) GetByTenantIDAndLLMName(tenantID string, llmName string // These factories append "___FactoryName" to the model name if factory == "LocalAI" || factory == "HuggingFace" || factory == "OpenAI-API-Compatible" { specialModelName := modelName + "___" + factory - err = DB.Where("tenant_id = ? AND llm_name = ?", tenantID, specialModelName).First(&tenantLLM).Error + err = db.WithContext(ctx).Where("tenant_id = ? AND llm_name = ?", tenantID, specialModelName).First(&tenantLLM).Error if err == nil { return &tenantLLM, nil } @@ -284,9 +287,9 @@ func (dao *TenantLLMDAO) GetByTenantIDAndLLMName(tenantID string, llmName string // Example: // // tenantLLM, err := dao.GetByTenantIDLLMNameAndFactory("tenant123", "gpt-4", "OpenAI") -func (dao *TenantLLMDAO) GetByTenantIDLLMNameAndFactory(tenantID, llmName, factory string) (*entity.TenantLLM, error) { +func (dao *TenantLLMDAO) GetByTenantIDLLMNameAndFactory(ctx context.Context, db *gorm.DB, tenantID, llmName, factory string) (*entity.TenantLLM, error) { var tenantLLM entity.TenantLLM - err := DB.Where("tenant_id = ? AND llm_name = ? AND llm_factory = ?", tenantID, llmName, factory).First(&tenantLLM).Error + err := db.WithContext(ctx).Where("tenant_id = ? AND llm_name = ? AND llm_factory = ?", tenantID, llmName, factory).First(&tenantLLM).Error if err != nil { return nil, err } @@ -294,8 +297,8 @@ func (dao *TenantLLMDAO) GetByTenantIDLLMNameAndFactory(tenantID, llmName, facto } // LookupTenantLLMByID looks up a TenantLLM record by ID and returns the record plus composite model name. -func LookupTenantLLMByID(tenantLLMDao *TenantLLMDAO, id int64) (*entity.TenantLLM, string, error) { - tenantLLM, err := tenantLLMDao.GetByID(id) +func LookupTenantLLMByID(ctx context.Context, db *gorm.DB, tenantLLMDao *TenantLLMDAO, id int64) (*entity.TenantLLM, string, error) { + tenantLLM, err := tenantLLMDao.GetByID(ctx, db, id) if err != nil { return nil, "", fmt.Errorf("failed to get tenant_llm by id %d: %w", id, err) } @@ -307,16 +310,16 @@ func LookupTenantLLMByID(tenantLLMDao *TenantLLMDAO, id int64) (*entity.TenantLL } // LookupTenantLLMByName looks up a TenantLLM record by tenant name and model type. -func LookupTenantLLMByName(tenantLLMDao *TenantLLMDAO, tenantID, name string, modelType entity.ModelType) (*entity.TenantLLM, string, error) { +func LookupTenantLLMByName(ctx context.Context, db *gorm.DB, tenantLLMDao *TenantLLMDAO, tenantID, name string, modelType entity.ModelType) (*entity.TenantLLM, string, error) { // Parse factory from name if present (e.g., "model@Factory") - modelName, factory := splitModelNameAndFactory(name) + modelName, factory := splitModelNameAndFactory(ctx, db, name) // If factory is found, use factory-based lookup if factory != "" { - return LookupTenantLLMByFactory(tenantLLMDao, tenantID, factory, modelName, modelType) + return LookupTenantLLMByFactory(ctx, db, tenantLLMDao, tenantID, factory, modelName, modelType) } - tenantLLM, err := tenantLLMDao.GetByTenantNameAndType(tenantID, modelName, modelType) + tenantLLM, err := tenantLLMDao.GetByTenantNameAndType(ctx, db, tenantID, modelName, modelType) if err != nil { return nil, "", fmt.Errorf("failed to get tenant_llm by name %s: %w", name, err) } @@ -328,8 +331,8 @@ func LookupTenantLLMByName(tenantLLMDao *TenantLLMDAO, tenantID, name string, mo } // LookupTenantLLMByFactory looks up a TenantLLM record by tenant, factory, and model name. -func LookupTenantLLMByFactory(tenantLLMDao *TenantLLMDAO, tenantID, factory, name string, modelType entity.ModelType) (*entity.TenantLLM, string, error) { - tenantLLM, err := tenantLLMDao.GetByTenantFactoryAndModelName(tenantID, factory, name) +func LookupTenantLLMByFactory(ctx context.Context, db *gorm.DB, tenantLLMDao *TenantLLMDAO, tenantID, factory, name string, modelType entity.ModelType) (*entity.TenantLLM, string, error) { + tenantLLM, err := tenantLLMDao.GetByTenantFactoryAndModelName(ctx, db, tenantID, factory, name) if err != nil { return nil, "", fmt.Errorf("failed to get tenant_llm by factory %s and name %s: %w", factory, name, err) } diff --git a/internal/dao/tenant_model.go b/internal/dao/tenant_model.go index ae2085e8e2..6e8c17ac67 100644 --- a/internal/dao/tenant_model.go +++ b/internal/dao/tenant_model.go @@ -17,6 +17,7 @@ package dao import ( + "context" "ragflow/internal/entity" "gorm.io/gorm" @@ -30,18 +31,18 @@ func NewTenantModelDAO() *TenantModelDAO { return &TenantModelDAO{} } -func (dao *TenantModelDAO) Create(instance *entity.TenantModel) error { - return DB.Create(instance).Error +func (dao *TenantModelDAO) Create(ctx context.Context, db *gorm.DB, instance *entity.TenantModel) error { + return db.WithContext(ctx).Create(instance).Error } -func (dao *TenantModelDAO) CreateBatch(models []*entity.TenantModel) error { +func (dao *TenantModelDAO) CreateBatch(ctx context.Context, db *gorm.DB, models []*entity.TenantModel) error { if len(models) == 0 { return nil } - return DB.Transaction(func(tx *gorm.DB) error { + return db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { for _, model := range models { - if err := tx.Create(model).Error; err != nil { + if err := tx.WithContext(ctx).Create(model).Error; err != nil { return err } } @@ -49,64 +50,64 @@ func (dao *TenantModelDAO) CreateBatch(models []*entity.TenantModel) error { }) } -func (dao *TenantModelDAO) DeleteByModelID(modelID string) (int64, error) { - result := DB.Unscoped().Where("id = ?", modelID).Delete(&entity.TenantModel{}) +func (dao *TenantModelDAO) DeleteByModelID(ctx context.Context, db *gorm.DB, modelID string) (int64, error) { + result := db.WithContext(ctx).Unscoped().Where("id = ?", modelID).Delete(&entity.TenantModel{}) return result.RowsAffected, result.Error } -func (dao *TenantModelDAO) DeleteByModelIDAndProviderIDAndInstanceID(modelID, providerID, instanceID string) (int64, error) { - result := DB.Unscoped().Where("id = ? AND provider_id = ? AND instance_id = ?", modelID, providerID, instanceID).Delete(&entity.TenantModel{}) +func (dao *TenantModelDAO) DeleteByModelIDAndProviderIDAndInstanceID(ctx context.Context, db *gorm.DB, modelID, providerID, instanceID string) (int64, error) { + result := db.WithContext(ctx).Unscoped().Where("id = ? AND provider_id = ? AND instance_id = ?", modelID, providerID, instanceID).Delete(&entity.TenantModel{}) return result.RowsAffected, result.Error } -func (dao *TenantModelDAO) DeleteByProviderIDAndInstanceID(provideID, instanceID string) (int64, error) { - result := DB.Unscoped().Where("provider_id = ? AND instance_id = ?", provideID, instanceID).Delete(&entity.TenantModel{}) +func (dao *TenantModelDAO) DeleteByProviderIDAndInstanceID(ctx context.Context, db *gorm.DB, provideID, instanceID string) (int64, error) { + result := db.WithContext(ctx).Unscoped().Where("provider_id = ? AND instance_id = ?", provideID, instanceID).Delete(&entity.TenantModel{}) return result.RowsAffected, result.Error } -func (dao *TenantModelDAO) DeleteByProviderIDAndInstanceIDAndModelName(provideID, instanceID, modelName string) (int64, error) { - result := DB.Unscoped().Where("provider_id = ? AND instance_id = ? AND model_name = ?", provideID, instanceID, modelName).Delete(&entity.TenantModel{}) +func (dao *TenantModelDAO) DeleteByProviderIDAndInstanceIDAndModelName(ctx context.Context, db *gorm.DB, provideID, instanceID, modelName string) (int64, error) { + result := db.WithContext(ctx).Unscoped().Where("provider_id = ? AND instance_id = ? AND model_name = ?", provideID, instanceID, modelName).Delete(&entity.TenantModel{}) return result.RowsAffected, result.Error } -func (dao *TenantModelDAO) UpdateStatusByIDAndScope(modelID, providerID, instanceID, status string) (int64, error) { - result := DB.Model(&entity.TenantModel{}).Where("id = ? AND provider_id = ? AND instance_id = ?", modelID, providerID, instanceID).Update("status", status) +func (dao *TenantModelDAO) UpdateStatusByIDAndScope(ctx context.Context, db *gorm.DB, modelID, providerID, instanceID, status string) (int64, error) { + result := db.WithContext(ctx).Model(&entity.TenantModel{}).Where("id = ? AND provider_id = ? AND instance_id = ?", modelID, providerID, instanceID).Update("status", status) return result.RowsAffected, result.Error } // GetByID get tenant model by primary key (id) -func (dao *TenantModelDAO) GetByID(id string) (*entity.TenantModel, error) { +func (dao *TenantModelDAO) GetByID(ctx context.Context, db *gorm.DB, id string) (*entity.TenantModel, error) { var model entity.TenantModel - err := DB.Where("id = ?", id).First(&model).Error + err := db.WithContext(ctx).Where("id = ?", id).First(&model).Error if err != nil { return nil, err } return &model, nil } -func (dao *TenantModelDAO) GetModelByProviderIDAndInstanceIDAndModelName(providerID, instanceID, modelName string) (*entity.TenantModel, error) { +func (dao *TenantModelDAO) GetModelByProviderIDAndInstanceIDAndModelName(ctx context.Context, db *gorm.DB, providerID, instanceID, modelName string) (*entity.TenantModel, error) { var model entity.TenantModel - err := DB.Where("provider_id = ? AND instance_id = ? AND model_name = ?", providerID, instanceID, modelName).First(&model).Error + err := db.WithContext(ctx).Where("provider_id = ? AND instance_id = ? AND model_name = ?", providerID, instanceID, modelName).First(&model).Error if err != nil { return nil, err } return &model, nil } -func (dao *TenantModelDAO) GetModelsByProviderIDAndInstanceIDAndModelName(providerID, instanceID, modelName string) ([]*entity.TenantModel, error) { +func (dao *TenantModelDAO) GetModelsByProviderIDAndInstanceIDAndModelName(ctx context.Context, db *gorm.DB, providerID, instanceID, modelName string) ([]*entity.TenantModel, error) { var models []*entity.TenantModel - err := DB.Where("provider_id = ? AND instance_id = ? AND model_name = ?", providerID, instanceID, modelName).Find(&models).Error + err := db.WithContext(ctx).Where("provider_id = ? AND instance_id = ? AND model_name = ?", providerID, instanceID, modelName).Find(&models).Error if err != nil { return nil, err } return models, nil } -func (dao *TenantModelDAO) GetByProviderIDAndInstanceIDAndModelTypeAndModelName(providerID, instanceID string, modelType int, modelName string) (*entity.TenantModel, error) { +func (dao *TenantModelDAO) GetByProviderIDAndInstanceIDAndModelTypeAndModelName(ctx context.Context, db *gorm.DB, providerID, instanceID string, modelType int, modelName string) (*entity.TenantModel, error) { var model entity.TenantModel // Use bitwise AND to match Python's bin_and(model_type) > 0 pattern. // A model_type value of 0 (unknown type) matches no row. - err := DB.Where("provider_id = ? AND instance_id = ? AND model_type & ? > 0 AND model_name = ?", providerID, instanceID, modelType, modelName).First(&model).Error + err := db.WithContext(ctx).Where("provider_id = ? AND instance_id = ? AND model_type & ? > 0 AND model_name = ?", providerID, instanceID, modelType, modelName).First(&model).Error if err != nil { return nil, err } @@ -114,9 +115,9 @@ func (dao *TenantModelDAO) GetByProviderIDAndInstanceIDAndModelTypeAndModelName( } // GetModelsByInstanceID get all models by instance ID -func (dao *TenantModelDAO) GetModelsByInstanceID(instanceID string) ([]*entity.TenantModel, error) { +func (dao *TenantModelDAO) GetModelsByInstanceID(ctx context.Context, db *gorm.DB, instanceID string) ([]*entity.TenantModel, error) { var models []*entity.TenantModel - err := DB.Where("instance_id = ?", instanceID).Find(&models).Error + err := db.WithContext(ctx).Where("instance_id = ?", instanceID).Find(&models).Error if err != nil { return nil, err } @@ -125,26 +126,26 @@ func (dao *TenantModelDAO) GetModelsByInstanceID(instanceID string) ([]*entity.T // DeleteByIDs deletes all models whose id is in the given list. // Mirrors Python's TenantModelService.delete_by_ids. -func (dao *TenantModelDAO) DeleteByIDs(ids []string) (int64, error) { +func (dao *TenantModelDAO) DeleteByIDs(ctx context.Context, db *gorm.DB, ids []string) (int64, error) { if len(ids) == 0 { return 0, nil } - result := DB.Unscoped().Where("id IN ?", ids).Delete(&entity.TenantModel{}) + result := db.WithContext(ctx).Unscoped().Where("id IN ?", ids).Delete(&entity.TenantModel{}) return result.RowsAffected, result.Error } // UpdateByID updates a tenant model's model_type and extra by primary key. // Mirrors Python's TenantModelService.update_model. -func (dao *TenantModelDAO) UpdateByID(id string, updates map[string]interface{}) error { - return DB.Model(&entity.TenantModel{}).Where("id = ?", id).Updates(updates).Error +func (dao *TenantModelDAO) UpdateByID(ctx context.Context, db *gorm.DB, id string, updates map[string]interface{}) error { + return db.WithContext(ctx).Model(&entity.TenantModel{}).Where("id = ?", id).Updates(updates).Error } // DeleteByInstanceIDs deletes all models whose instance_id is in the given list. -func (dao *TenantModelDAO) DeleteByInstanceIDs(instanceIDs []string) (int64, error) { +func (dao *TenantModelDAO) DeleteByInstanceIDs(ctx context.Context, db *gorm.DB, instanceIDs []string) (int64, error) { if len(instanceIDs) == 0 { return 0, nil } - result := DB.Unscoped().Where("instance_id IN ?", instanceIDs).Delete(&entity.TenantModel{}) + result := db.WithContext(ctx).Unscoped().Where("instance_id IN ?", instanceIDs).Delete(&entity.TenantModel{}) return result.RowsAffected, result.Error } @@ -156,12 +157,12 @@ func (dao *TenantModelDAO) DeleteByInstanceIDs(instanceIDs []string) (int64, err // /api/v1/models response assembly. The Go port never WRITES to // tenant_model, so callers must treat an empty result as "use factory // defaults" — see ModelProviderService.ListTenantAddedModels. -func (dao *TenantModelDAO) GetModelsByProviderIDsAndInstanceIDs(providerIDs, instanceIDs []string) ([]*entity.TenantModel, error) { +func (dao *TenantModelDAO) GetModelsByProviderIDsAndInstanceIDs(ctx context.Context, db *gorm.DB, providerIDs, instanceIDs []string) ([]*entity.TenantModel, error) { models := make([]*entity.TenantModel, 0) if len(providerIDs) == 0 || len(instanceIDs) == 0 { return models, nil } - err := DB.Where("provider_id IN ? AND instance_id IN ?", providerIDs, instanceIDs).Find(&models).Error + err := db.WithContext(ctx).Where("provider_id IN ? AND instance_id IN ?", providerIDs, instanceIDs).Find(&models).Error if err != nil { return nil, err } diff --git a/internal/dao/tenant_model_group.go b/internal/dao/tenant_model_group.go index e2d26982c9..4a86b5bea6 100644 --- a/internal/dao/tenant_model_group.go +++ b/internal/dao/tenant_model_group.go @@ -17,7 +17,10 @@ package dao import ( + "context" "ragflow/internal/entity" + + "gorm.io/gorm" ) // TenantModelGroupDAO tenant model group data access object @@ -29,9 +32,9 @@ func NewTenantModelGroupDAO() *TenantModelGroupDAO { } // GetByID get tenant model group by primary key (id) -func (dao *TenantModelGroupDAO) GetByID(id string) (*entity.TenantModelGroup, error) { +func (dao *TenantModelGroupDAO) GetByID(ctx context.Context, db *gorm.DB, id string) (*entity.TenantModelGroup, error) { var group entity.TenantModelGroup - err := DB.Where("id = ?", id).First(&group).Error + err := db.WithContext(ctx).Where("id = ?", id).First(&group).Error if err != nil { return nil, err } diff --git a/internal/dao/tenant_model_group_mapping.go b/internal/dao/tenant_model_group_mapping.go index c06270d275..f51a5a595f 100644 --- a/internal/dao/tenant_model_group_mapping.go +++ b/internal/dao/tenant_model_group_mapping.go @@ -17,7 +17,10 @@ package dao import ( + "context" "ragflow/internal/entity" + + "gorm.io/gorm" ) // TenantModelGroupMappingDAO tenant model group mapping data access object @@ -29,9 +32,9 @@ func NewTenantModelGroupMappingDAO() *TenantModelGroupMappingDAO { } // GetByID get tenant model group mapping by composite primary key -func (dao *TenantModelGroupMappingDAO) GetByID(groupID, providerID, instanceID, modelID string) (*entity.TenantModelGroupMapping, error) { +func (dao *TenantModelGroupMappingDAO) GetByID(ctx context.Context, db *gorm.DB, groupID, providerID, instanceID, modelID string) (*entity.TenantModelGroupMapping, error) { var mapping entity.TenantModelGroupMapping - err := DB.Where("group_id = ? AND provider_id = ? AND instance_id = ? AND model_id = ?", groupID, providerID, instanceID, modelID).First(&mapping).Error + err := db.WithContext(ctx).Where("group_id = ? AND provider_id = ? AND instance_id = ? AND model_id = ?", groupID, providerID, instanceID, modelID).First(&mapping).Error if err != nil { return nil, err } diff --git a/internal/dao/tenant_model_instance.go b/internal/dao/tenant_model_instance.go index 41cc4688f9..0a0cfc70f7 100644 --- a/internal/dao/tenant_model_instance.go +++ b/internal/dao/tenant_model_instance.go @@ -17,6 +17,7 @@ package dao import ( + "context" "errors" "fmt" "ragflow/internal/entity" @@ -32,19 +33,19 @@ func NewTenantModelInstanceDAO() *TenantModelInstanceDAO { return &TenantModelInstanceDAO{} } -func (dao *TenantModelInstanceDAO) Create(instance *entity.TenantModelInstance) error { +func (dao *TenantModelInstanceDAO) Create(ctx context.Context, db *gorm.DB, instance *entity.TenantModelInstance) error { // begin tx and check if the same provider instance exists - tx := DB.Begin() + tx := db.WithContext(ctx).Begin() defer tx.Rollback() var existingInstance entity.TenantModelInstance - err := tx.Where("provider_id = ? AND instance_name = ?", instance.ProviderID, instance.InstanceName).First(&existingInstance).Error + err := tx.WithContext(ctx).Where("provider_id = ? AND instance_name = ?", instance.ProviderID, instance.InstanceName).First(&existingInstance).Error if err == nil { return fmt.Errorf("instance %s already exists", instance.InstanceName) } if !errors.Is(err, gorm.ErrRecordNotFound) { return err } - err = tx.Create(instance).Error + err = tx.WithContext(ctx).Create(instance).Error if err != nil { return err } @@ -52,9 +53,9 @@ func (dao *TenantModelInstanceDAO) Create(instance *entity.TenantModelInstance) return nil } -func (dao *TenantModelInstanceDAO) GetAllInstancesByProviderID(providerID string) ([]*entity.TenantModelInstance, error) { +func (dao *TenantModelInstanceDAO) GetAllInstancesByProviderID(ctx context.Context, db *gorm.DB, providerID string) ([]*entity.TenantModelInstance, error) { var instances []*entity.TenantModelInstance - err := DB.Where("provider_id = ?", providerID).Find(&instances).Error + err := db.WithContext(ctx).Where("provider_id = ?", providerID).Find(&instances).Error if err != nil { return nil, err } @@ -66,30 +67,30 @@ func (dao *TenantModelInstanceDAO) GetAllInstancesByProviderID(providerID string // TenantModelInstanceService.get_by_provider_ids used by // models_api_service.list_tenant_added_models. An empty input slice // returns an empty (non-nil) slice with no error. -func (dao *TenantModelInstanceDAO) GetByProviderIDs(providerIDs []string) ([]*entity.TenantModelInstance, error) { +func (dao *TenantModelInstanceDAO) GetByProviderIDs(ctx context.Context, db *gorm.DB, providerIDs []string) ([]*entity.TenantModelInstance, error) { instances := make([]*entity.TenantModelInstance, 0) if len(providerIDs) == 0 { return instances, nil } - err := DB.Where("provider_id IN ?", providerIDs).Find(&instances).Error + err := db.WithContext(ctx).Where("provider_id IN ?", providerIDs).Find(&instances).Error if err != nil { return nil, err } return instances, nil } -func (dao *TenantModelInstanceDAO) GetInstanceByApiKey(apiKey, providerID string) (*entity.TenantModelInstance, error) { +func (dao *TenantModelInstanceDAO) GetInstanceByApiKey(ctx context.Context, db *gorm.DB, apiKey, providerID string) (*entity.TenantModelInstance, error) { var instance entity.TenantModelInstance - err := DB.Where("api_key = ? && provider_id = ?", apiKey, providerID).First(&instance).Error + err := db.WithContext(ctx).Where("api_key = ? AND provider_id = ?", apiKey, providerID).First(&instance).Error if err != nil { return nil, err } return &instance, nil } -func (dao *TenantModelInstanceDAO) GetByProviderIDAndInstanceName(providerID, instanceName string) (*entity.TenantModelInstance, error) { +func (dao *TenantModelInstanceDAO) GetByProviderIDAndInstanceName(ctx context.Context, db *gorm.DB, providerID, instanceName string) (*entity.TenantModelInstance, error) { var instance entity.TenantModelInstance - err := DB.Where("provider_id = ? AND instance_name = ?", providerID, instanceName).First(&instance).Error + err := db.WithContext(ctx).Where("provider_id = ? AND instance_name = ?", providerID, instanceName).First(&instance).Error if err != nil { return nil, err } @@ -97,38 +98,38 @@ func (dao *TenantModelInstanceDAO) GetByProviderIDAndInstanceName(providerID, in } // GetByID get tenant model instance by primary key (id) -func (dao *TenantModelInstanceDAO) GetByID(id string) (*entity.TenantModelInstance, error) { +func (dao *TenantModelInstanceDAO) GetByID(ctx context.Context, db *gorm.DB, id string) (*entity.TenantModelInstance, error) { var instance entity.TenantModelInstance - err := DB.Where("id = ?", id).First(&instance).Error + err := db.WithContext(ctx).Where("id = ?", id).First(&instance).Error if err != nil { return nil, err } return &instance, nil } -func (dao *TenantModelInstanceDAO) DeleteByProviderIDAndInstanceName(providerID, instanceName string) (int64, error) { - result := DB.Unscoped().Where("provider_id = ? and instance_name = ?", providerID, instanceName).Delete(&entity.TenantModelInstance{}) +func (dao *TenantModelInstanceDAO) DeleteByProviderIDAndInstanceName(ctx context.Context, db *gorm.DB, providerID, instanceName string) (int64, error) { + result := db.WithContext(ctx).Unscoped().Where("provider_id = ? and instance_name = ?", providerID, instanceName).Delete(&entity.TenantModelInstance{}) return result.RowsAffected, result.Error } // UpdateByID updates a tenant model instance by primary key. // Mirrors Python's TenantModelInstanceService.update_by_id. -func (dao *TenantModelInstanceDAO) UpdateByID(id string, updates map[string]interface{}) error { - return DB.Model(&entity.TenantModelInstance{}).Where("id = ?", id).Updates(updates).Error +func (dao *TenantModelInstanceDAO) UpdateByID(ctx context.Context, db *gorm.DB, id string, updates map[string]interface{}) error { + return db.WithContext(ctx).Model(&entity.TenantModelInstance{}).Where("id = ?", id).Updates(updates).Error } // DeleteByIDs deletes all instances whose id is in the given list. // Mirrors Python's TenantModelInstanceService.delete_by_ids. -func (dao *TenantModelInstanceDAO) DeleteByIDs(ids []string) (int64, error) { +func (dao *TenantModelInstanceDAO) DeleteByIDs(ctx context.Context, db *gorm.DB, ids []string) (int64, error) { if len(ids) == 0 { return 0, nil } - result := DB.Unscoped().Where("id IN ?", ids).Delete(&entity.TenantModelInstance{}) + result := db.WithContext(ctx).Unscoped().Where("id IN ?", ids).Delete(&entity.TenantModelInstance{}) return result.RowsAffected, result.Error } // DeleteByProviderID deletes all instances for the given provider. -func (dao *TenantModelInstanceDAO) DeleteByProviderID(providerID string) (int64, error) { - result := DB.Unscoped().Where("provider_id = ?", providerID).Delete(&entity.TenantModelInstance{}) +func (dao *TenantModelInstanceDAO) DeleteByProviderID(ctx context.Context, db *gorm.DB, providerID string) (int64, error) { + result := db.WithContext(ctx).Unscoped().Where("provider_id = ?", providerID).Delete(&entity.TenantModelInstance{}) return result.RowsAffected, result.Error } diff --git a/internal/dao/tenant_model_provider.go b/internal/dao/tenant_model_provider.go index 58827ecaf7..857fbe5750 100644 --- a/internal/dao/tenant_model_provider.go +++ b/internal/dao/tenant_model_provider.go @@ -17,7 +17,10 @@ package dao import ( + "context" "ragflow/internal/entity" + + "gorm.io/gorm" ) // TenantModelProviderDAO tenant model provider data access object @@ -28,14 +31,14 @@ func NewTenantModelProviderDAO() *TenantModelProviderDAO { return &TenantModelProviderDAO{} } -func (dao *TenantModelProviderDAO) Create(provider *entity.TenantModelProvider) error { - return DB.Create(provider).Error +func (dao *TenantModelProviderDAO) Create(ctx context.Context, db *gorm.DB, provider *entity.TenantModelProvider) error { + return db.WithContext(ctx).Create(provider).Error } // GetByID get tenant model provider by primary key (id) -func (dao *TenantModelProviderDAO) GetByID(id string) (*entity.TenantModelProvider, error) { +func (dao *TenantModelProviderDAO) GetByID(ctx context.Context, db *gorm.DB, id string) (*entity.TenantModelProvider, error) { var provider entity.TenantModelProvider - err := DB.Where("id = ?", id).First(&provider).Error + err := db.WithContext(ctx).Where("id = ?", id).First(&provider).Error if err != nil { return nil, err } @@ -43,9 +46,9 @@ func (dao *TenantModelProviderDAO) GetByID(id string) (*entity.TenantModelProvid } // GetByTenantIDAndProviderName get the providers by tenant ID and provider name -func (dao *TenantModelProviderDAO) GetByTenantIDAndProviderName(tenantID, providerName string) (*entity.TenantModelProvider, error) { +func (dao *TenantModelProviderDAO) GetByTenantIDAndProviderName(ctx context.Context, db *gorm.DB, tenantID, providerName string) (*entity.TenantModelProvider, error) { var provider entity.TenantModelProvider - err := DB.Where("tenant_id = ? AND provider_name = ?", tenantID, providerName).First(&provider).Error + err := db.WithContext(ctx).Where("tenant_id = ? AND provider_name = ?", tenantID, providerName).First(&provider).Error if err != nil { return nil, err } @@ -53,21 +56,21 @@ func (dao *TenantModelProviderDAO) GetByTenantIDAndProviderName(tenantID, provid } // DeleteByTenantID deletes all model providers by tenant ID (hard delete) -func (dao *TenantModelProviderDAO) DeleteByTenantID(tenantID string) (int64, error) { - result := DB.Unscoped().Where("tenant_id = ?", tenantID).Delete(&entity.TenantModelProvider{}) +func (dao *TenantModelProviderDAO) DeleteByTenantID(ctx context.Context, db *gorm.DB, tenantID string) (int64, error) { + result := db.WithContext(ctx).Unscoped().Where("tenant_id = ?", tenantID).Delete(&entity.TenantModelProvider{}) return result.RowsAffected, result.Error } -// DeleteByTenantID deletes all providers by tenant ID (hard delete) -func (dao *TenantModelProviderDAO) DeleteByTenantIDAndProviderName(tenantID, providerName string) (int64, error) { - result := DB.Unscoped().Where("tenant_id = ? AND provider_name = ?", tenantID, providerName).Delete(&entity.TenantModelProvider{}) +// DeleteByTenantIDAndProviderName DeleteByTenantID deletes all providers by tenant ID (hard delete) +func (dao *TenantModelProviderDAO) DeleteByTenantIDAndProviderName(ctx context.Context, db *gorm.DB, tenantID, providerName string) (int64, error) { + result := db.WithContext(ctx).Unscoped().Where("tenant_id = ? AND provider_name = ?", tenantID, providerName).Delete(&entity.TenantModelProvider{}) return result.RowsAffected, result.Error } // ListByID list tenant model providers by ID -func (dao *TenantModelProviderDAO) ListByID(id string) ([]string, error) { +func (dao *TenantModelProviderDAO) ListByID(ctx context.Context, db *gorm.DB, id string) ([]string, error) { var providerNames []string - err := DB.Model(&entity.TenantModelProvider{}). + err := db.WithContext(ctx).Model(&entity.TenantModelProvider{}). Where("tenant_id = ?", id). Pluck("provider_name", &providerNames).Error return providerNames, err @@ -79,9 +82,9 @@ func (dao *TenantModelProviderDAO) ListByID(id string) ([]string, error) { // uses this to enumerate which providers a tenant has linked before // fanning out to TenantModelInstanceDAO / TenantModelDAO for the joined // result. -func (dao *TenantModelProviderDAO) GetByTenantID(tenantID string) ([]*entity.TenantModelProvider, error) { +func (dao *TenantModelProviderDAO) GetByTenantID(ctx context.Context, db *gorm.DB, tenantID string) ([]*entity.TenantModelProvider, error) { var providers []*entity.TenantModelProvider - err := DB.Where("tenant_id = ?", tenantID).Find(&providers).Error + err := db.WithContext(ctx).Where("tenant_id = ?", tenantID).Find(&providers).Error if err != nil { return nil, err } diff --git a/internal/dao/tenant_model_test.go b/internal/dao/tenant_model_test.go index 45f1125c44..deff4bd9c9 100644 --- a/internal/dao/tenant_model_test.go +++ b/internal/dao/tenant_model_test.go @@ -17,6 +17,7 @@ package dao import ( + "context" "testing" "github.com/glebarez/sqlite" @@ -31,7 +32,7 @@ func setupTenantModelDAOTestDB(t *testing.T) *gorm.DB { if err != nil { t.Fatalf("failed to open sqlite: %v", err) } - if err := db.AutoMigrate(&entity.TenantModel{}); err != nil { + if err = db.AutoMigrate(&entity.TenantModel{}); err != nil { t.Fatalf("failed to migrate tenant_model: %v", err) } return db @@ -58,7 +59,8 @@ func TestTenantModelDAODeleteByModelIDAndScopeDeletesOnlyMatchingModel(t *testin seedTenantModel(t, db, &entity.TenantModel{ID: "model-delete", ModelName: "m", ModelType: int(entity.ModelTypeChat), ProviderID: "provider-1", InstanceID: "instance-1", Status: "active"}) seedTenantModel(t, db, &entity.TenantModel{ID: "model-keep", ModelName: "m", ModelType: int(entity.ModelTypeChat), ProviderID: "provider-1", InstanceID: "instance-2", Status: "active"}) - rows, err := NewTenantModelDAO().DeleteByModelIDAndProviderIDAndInstanceID("model-delete", "provider-1", "instance-1") + ctx := context.Background() + rows, err := NewTenantModelDAO().DeleteByModelIDAndProviderIDAndInstanceID(ctx, db, "model-delete", "provider-1", "instance-1") if err != nil { t.Fatalf("DeleteByModelIDAndProviderIDAndInstanceID() error = %v", err) } @@ -67,13 +69,13 @@ func TestTenantModelDAODeleteByModelIDAndScopeDeletesOnlyMatchingModel(t *testin } var count int64 - if err := db.Model(&entity.TenantModel{}).Where("id = ?", "model-delete").Count(&count).Error; err != nil { + if err = db.Model(&entity.TenantModel{}).Where("id = ?", "model-delete").Count(&count).Error; err != nil { t.Fatalf("count deleted model: %v", err) } if count != 0 { t.Fatalf("deleted model count = %d, want 0", count) } - if err := db.Model(&entity.TenantModel{}).Where("id = ?", "model-keep").Count(&count).Error; err != nil { + if err = db.Model(&entity.TenantModel{}).Where("id = ?", "model-keep").Count(&count).Error; err != nil { t.Fatalf("count kept model: %v", err) } if count != 1 { @@ -87,7 +89,8 @@ func TestTenantModelDAOUpdateStatusByIDAndScope(t *testing.T) { seedTenantModel(t, db, &entity.TenantModel{ID: "model-status", ModelName: "m", ModelType: int(entity.ModelTypeChat), ProviderID: "provider-1", InstanceID: "instance-1", Status: "active"}) - rows, err := NewTenantModelDAO().UpdateStatusByIDAndScope("model-status", "provider-1", "instance-1", "inactive") + ctx := context.Background() + rows, err := NewTenantModelDAO().UpdateStatusByIDAndScope(ctx, db, "model-status", "provider-1", "instance-1", "inactive") if err != nil { t.Fatalf("UpdateStatusByIDAndScope() error = %v", err) } @@ -96,14 +99,14 @@ func TestTenantModelDAOUpdateStatusByIDAndScope(t *testing.T) { } var got entity.TenantModel - if err := db.Where("id = ?", "model-status").First(&got).Error; err != nil { + if err = db.Where("id = ?", "model-status").First(&got).Error; err != nil { t.Fatalf("failed to reload model: %v", err) } if got.Status != "inactive" { t.Fatalf("status = %q, want inactive", got.Status) } - rows, err = NewTenantModelDAO().UpdateStatusByIDAndScope("model-status", "provider-1", "wrong-instance", "active") + rows, err = NewTenantModelDAO().UpdateStatusByIDAndScope(ctx, db, "model-status", "provider-1", "wrong-instance", "active") if err != nil { t.Fatalf("UpdateStatusByIDAndScope() wrong scope error = %v", err) } diff --git a/internal/dao/time_record.go b/internal/dao/time_record.go index 06d532ec4d..f45747255e 100644 --- a/internal/dao/time_record.go +++ b/internal/dao/time_record.go @@ -17,7 +17,10 @@ package dao import ( + "context" "ragflow/internal/entity" + + "gorm.io/gorm" ) // TimeRecordDAO time record data access object @@ -29,14 +32,14 @@ func NewTimeRecordDAO() *TimeRecordDAO { } // Create inserts a new record -func (dao *TimeRecordDAO) Create(record *entity.TimeRecord) error { - return DB.Create(record).Error +func (dao *TimeRecordDAO) Create(ctx context.Context, db *gorm.DB, record *entity.TimeRecord) error { + return db.WithContext(ctx).Create(record).Error } // GetRecent retrieves the most recently inserted records (ordered by ID descending) -func (dao *TimeRecordDAO) GetRecent(limit int) ([]*entity.TimeRecord, error) { +func (dao *TimeRecordDAO) GetRecent(ctx context.Context, db *gorm.DB, limit int) ([]*entity.TimeRecord, error) { var records []*entity.TimeRecord - err := DB.Order("id DESC").Limit(limit).Find(&records).Error + err := db.WithContext(ctx).Order("id DESC").Limit(limit).Find(&records).Error if err != nil { return nil, err } @@ -44,21 +47,21 @@ func (dao *TimeRecordDAO) GetRecent(limit int) ([]*entity.TimeRecord, error) { } // GetCount returns the total number of records -func (dao *TimeRecordDAO) GetCount() (int64, error) { +func (dao *TimeRecordDAO) GetCount(ctx context.Context, db *gorm.DB) (int64, error) { var count int64 - err := DB.Model(&entity.TimeRecord{}).Count(&count).Error + err := db.WithContext(ctx).Model(&entity.TimeRecord{}).Count(&count).Error return count, err } // DeleteOldest removes the oldest records (smallest ID) with limit -func (dao *TimeRecordDAO) DeleteOldest(limit int64) error { - return DB.Exec("DELETE FROM time_records ORDER BY id ASC LIMIT ?", limit).Error +func (dao *TimeRecordDAO) DeleteOldest(ctx context.Context, db *gorm.DB, limit int64) error { + return db.WithContext(ctx).Exec("DELETE FROM time_records ORDER BY id ASC LIMIT ?", limit).Error } // GetByID retrieves a single record by its ID -func (dao *TimeRecordDAO) GetByID(id int64) (*entity.TimeRecord, error) { +func (dao *TimeRecordDAO) GetByID(ctx context.Context, db *gorm.DB, id int64) (*entity.TimeRecord, error) { var record entity.TimeRecord - err := DB.First(&record, id).Error + err := db.WithContext(ctx).First(&record, id).Error if err != nil { return nil, err } @@ -66,17 +69,17 @@ func (dao *TimeRecordDAO) GetByID(id int64) (*entity.TimeRecord, error) { } // GetAll retrieves all records -func (dao *TimeRecordDAO) GetAll() ([]*entity.TimeRecord, error) { +func (dao *TimeRecordDAO) GetAll(ctx context.Context, db *gorm.DB) ([]*entity.TimeRecord, error) { var records []*entity.TimeRecord - err := DB.Find(&records).Error + err := db.WithContext(ctx).Find(&records).Error return records, err } // KeepLatest keeps the latest N records and deletes older ones -func (dao *TimeRecordDAO) KeepLatest(count int64) error { +func (dao *TimeRecordDAO) KeepLatest(ctx context.Context, db *gorm.DB, count int64) error { // Step 1: Get the maximum ID var maxID int64 - if err := DB.Model(&entity.TimeRecord{}).Select("COALESCE(MAX(id), 0)").Scan(&maxID).Error; err != nil { + if err := db.WithContext(ctx).Model(&entity.TimeRecord{}).Select("COALESCE(MAX(id), 0)").Scan(&maxID).Error; err != nil { return err } @@ -94,10 +97,10 @@ func (dao *TimeRecordDAO) KeepLatest(count int64) error { } // Step 3: Delete records with ID <= threshold - return DB.Where("id <= ?", thresholdID).Delete(&entity.TimeRecord{}).Error + return db.WithContext(ctx).Where("id <= ?", thresholdID).Delete(&entity.TimeRecord{}).Error } // DeleteAll deletes all records -func (dao *TimeRecordDAO) DeleteAll() error { - return DB.Where("1=1").Delete(&entity.TimeRecord{}).Error +func (dao *TimeRecordDAO) DeleteAll(ctx context.Context, db *gorm.DB) error { + return db.WithContext(ctx).Where("1=1").Delete(&entity.TimeRecord{}).Error } diff --git a/internal/handler/providers.go b/internal/handler/providers.go index c11a4318ad..9fb4225231 100644 --- a/internal/handler/providers.go +++ b/internal/handler/providers.go @@ -74,9 +74,10 @@ func (h *ProviderHandler) ListProviders(c *gin.Context) { } userID := c.GetString("user_id") + ctx := c.Request.Context() // list tenant providers - providers, errorCode, err := h.modelProviderService.ListProvidersOfTenant(userID) + providers, errorCode, err := h.modelProviderService.ListProvidersOfTenant(ctx, userID) if err != nil { common.ResponseWithCodeData(c, errorCode, nil, err.Error()) return @@ -99,8 +100,9 @@ func (h *ProviderHandler) AddProvider(c *gin.Context) { } userID := c.GetString("user_id") + ctx := c.Request.Context() - errorCode, err := h.modelProviderService.AddModelProvider(req.ProviderName, userID) + errorCode, err := h.modelProviderService.AddModelProvider(ctx, req.ProviderName, userID) if err != nil { common.ErrorWithCode(c, errorCode, err.Error()) return @@ -117,8 +119,9 @@ func (h *ProviderHandler) DeleteProvider(c *gin.Context) { } userID := c.GetString("user_id") + ctx := c.Request.Context() - errorCode, err := h.modelProviderService.DeleteModelProvider(userID, providerName) + errorCode, err := h.modelProviderService.DeleteModelProvider(ctx, userID, providerName) if err != nil { common.ErrorWithCode(c, errorCode, err.Error()) return @@ -265,7 +268,7 @@ func (h *ProviderHandler) CreateProviderInstance(c *gin.Context) { // instance without API key validation or model creation. // Mirrors Python's provider_api.py:349 — set(data.keys()) == {"instance_name"}. if req.APIKey == "" && req.BaseURL == "" && req.Region == "" && len(req.ModelInfo) == 0 { - code, err := h.modelProviderService.CreateNameOnlyProviderInstance(providerName, req.InstanceName, userID) + code, err := h.modelProviderService.CreateNameOnlyProviderInstance(ctx, providerName, req.InstanceName, userID) if err != nil { common.ErrorWithCode(c, code, err.Error()) return @@ -291,8 +294,9 @@ func (h *ProviderHandler) ListProviderInstances(c *gin.Context) { } userID := c.GetString("user_id") + ctx := c.Request.Context() - instances, errorCode, err := h.modelProviderService.ListProviderInstances(providerName, userID) + instances, errorCode, err := h.modelProviderService.ListProviderInstances(ctx, providerName, userID) if err != nil { common.ErrorWithCode(c, errorCode, err.Error()) return @@ -315,8 +319,9 @@ func (h *ProviderHandler) ShowProviderInstance(c *gin.Context) { } userID := c.GetString("user_id") + ctx := c.Request.Context() - instance, errorCode, err := h.modelProviderService.ShowProviderInstance(providerName, instanceIDOrName, userID) + instance, errorCode, err := h.modelProviderService.ShowProviderInstance(ctx, providerName, instanceIDOrName, userID) if err != nil { common.ErrorWithCode(c, errorCode, err.Error()) return @@ -388,10 +393,9 @@ func (h *ProviderHandler) CheckInstanceConnection(c *gin.Context) { common.ResponseWithHttpCodeData(c, http.StatusBadRequest, 400, nil, "Instance name is required") return } - userID := c.GetString("user_id") - instanceInfo, code, err := h.modelProviderService.ShowProviderInstance(providerName, instanceName, userID) + instanceInfo, code, err := h.modelProviderService.ShowProviderInstance(ctx, providerName, instanceName, userID) if err != nil { common.ErrorWithCode(c, code, err.Error()) return @@ -535,8 +539,9 @@ func (h *ProviderHandler) DropProviderInstance(c *gin.Context) { } userID := c.GetString("user_id") + ctx := c.Request.Context() - code, err := h.modelProviderService.DropProviderInstances(providerName, userID, req.Instances) + code, err := h.modelProviderService.DropProviderInstances(ctx, providerName, userID, req.Instances) if err != nil { common.ErrorWithCode(c, code, err.Error()) return @@ -578,7 +583,7 @@ func (h *ProviderHandler) ListInstanceModels(c *gin.Context) { return } - modelInstances, err := h.modelProviderService.ListInstanceModels(providerName, instanceName, c.GetString("user_id")) + modelInstances, err := h.modelProviderService.ListInstanceModels(ctx, providerName, instanceName, c.GetString("user_id")) if err != nil { common.ErrorWithCode(c, common.CodeNotFound, err.Error()) return @@ -648,7 +653,8 @@ func (h *ProviderHandler) AlterModel(c *gin.Context) { return } - code, err := h.modelProviderService.AlterModel(providerName, instanceName, modelName, userID, modelID, updateDict) + ctx := c.Request.Context() + code, err := h.modelProviderService.AlterModel(ctx, providerName, instanceName, modelName, userID, modelID, updateDict) if err != nil { common.ErrorWithCode(c, code, err.Error()) return @@ -659,19 +665,19 @@ func (h *ProviderHandler) AlterModel(c *gin.Context) { func prepareProviderInstance(providerName, instanceName, reqProviderName, reqInstanceName string) error { if providerName == "" { - return errors.New("Provider name is required") + return errors.New("provider name is required") } if instanceName == "" { - return errors.New("Instance name is required") + return errors.New("instance name is required") } if reqProviderName != "" && !strings.EqualFold(reqProviderName, providerName) { - return errors.New("Provider name does not match path") + return errors.New("provider name does not match path") } if reqInstanceName != "" && !strings.EqualFold(reqInstanceName, instanceName) { - return errors.New("Instance name does not match path") + return errors.New("instance name does not match path") } return nil @@ -704,8 +710,9 @@ func (h *ProviderHandler) AddModel(c *gin.Context) { } userID := c.GetString("user_id") + ctx := c.Request.Context() - code, err := h.modelProviderService.AddModel(&req, userID) + code, err := h.modelProviderService.AddModel(ctx, &req, userID) if err != nil { common.ErrorWithCode(c, code, err.Error()) return @@ -741,8 +748,9 @@ func (h *ProviderHandler) DropInstanceModels(c *gin.Context) { } userID := c.GetString("user_id") + ctx := c.Request.Context() - code, err := h.modelProviderService.DropInstanceModels(providerName, instanceName, userID, req.ModelNames) + code, err := h.modelProviderService.DropInstanceModels(ctx, providerName, instanceName, userID, req.ModelNames) if err != nil { common.ErrorWithCode(c, code, err.Error()) return diff --git a/internal/ingestion/component/dispatch_model.go b/internal/ingestion/component/dispatch_model.go index 73d3f5d2a1..b6c74a4e76 100644 --- a/internal/ingestion/component/dispatch_model.go +++ b/internal/ingestion/component/dispatch_model.go @@ -67,12 +67,12 @@ func defaultResolveTenantModelByType(ctx context.Context, db *gorm.DB, tenantID return nil, "", nil, 0, fmt.Errorf("no default %s model is set", modelType) } if tenantModelID := tenantModelIDByType(tenant, modelType); tenantModelID != "" { - driver, modelName, apiConfig, maxTokens, err := resolveModelConfigByID(tenantID, modelType, tenantModelID) + driver, modelName, apiConfig, maxTokens, err := resolveModelConfigByID(ctx, db, tenantID, modelType, tenantModelID) if err == nil { return driver, modelName, apiConfig, maxTokens, nil } } - return resolveModelConfig(tenantID, modelType, modelID) + return resolveModelConfig(ctx, db, tenantID, modelType, modelID) } func tenantModelIDByType(tenant *entity.Tenant, modelType entity.ModelType) string { @@ -106,22 +106,22 @@ func stringValue(value *string) string { return *value } -func resolveModelConfig(tenantID string, modelType entity.ModelType, modelRef string) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { +func resolveModelConfig(ctx context.Context, db *gorm.DB, tenantID string, modelType entity.ModelType, modelRef string) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { modelDAO := dao.NewTenantModelDAO() - if _, err := modelDAO.GetByID(modelRef); err == nil { - return resolveModelConfigByID(tenantID, modelType, modelRef) + if _, err := modelDAO.GetByID(ctx, db, modelRef); err == nil { + return resolveModelConfigByID(ctx, db, tenantID, modelType, modelRef) } else if !errorsIsRecordNotFound(err) { return nil, "", nil, 0, err } - return resolveModelConfigFromProviderInstance(tenantID, modelType, modelRef) + return resolveModelConfigFromProviderInstance(ctx, db, tenantID, modelType, modelRef) } -func resolveModelConfigByID(tenantID string, modelType entity.ModelType, modelID string) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { +func resolveModelConfigByID(ctx context.Context, db *gorm.DB, tenantID string, modelType entity.ModelType, modelID string) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { modelDAO := dao.NewTenantModelDAO() instanceDAO := dao.NewTenantModelInstanceDAO() providerDAO := dao.NewTenantModelProviderDAO() - modelObj, err := modelDAO.GetByID(modelID) + modelObj, err := modelDAO.GetByID(ctx, db, modelID) if err != nil { return nil, "", nil, 0, err } @@ -131,11 +131,11 @@ func resolveModelConfigByID(tenantID string, modelType entity.ModelType, modelID if !entity.ModelType(modelObj.ModelType).Has(modelType) { return nil, "", nil, 0, fmt.Errorf("model %q cannot be used as %s model", modelID, modelType.String()) } - instance, err := instanceDAO.GetByID(modelObj.InstanceID) + instance, err := instanceDAO.GetByID(ctx, db, modelObj.InstanceID) if err != nil { return nil, "", nil, 0, err } - provider, err := providerDAO.GetByID(modelObj.ProviderID) + provider, err := providerDAO.GetByID(ctx, db, modelObj.ProviderID) if err != nil { return nil, "", nil, 0, err } @@ -174,7 +174,7 @@ func resolveModelConfigByID(tenantID string, modelType entity.ModelType, modelID return driver, modelObj.ModelName, apiConfig, maxTokens, nil } -func resolveModelConfigFromProviderInstance(tenantID string, modelType entity.ModelType, modelName string) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { +func resolveModelConfigFromProviderInstance(ctx context.Context, db *gorm.DB, tenantID string, modelType entity.ModelType, modelName string) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { pureModelName, instanceName, providerName, err := parseCompositeModelName(modelName) if err != nil { return nil, "", nil, 0, err @@ -184,11 +184,11 @@ func resolveModelConfigFromProviderInstance(tenantID string, modelType entity.Mo instanceDAO := dao.NewTenantModelInstanceDAO() modelDAO := dao.NewTenantModelDAO() - provider, err := providerDAO.GetByTenantIDAndProviderName(tenantID, providerName) + provider, err := providerDAO.GetByTenantIDAndProviderName(ctx, db, tenantID, providerName) if err != nil { return nil, "", nil, 0, fmt.Errorf("provider %q lookup failed: %w", providerName, err) } - instance, err := instanceDAO.GetByProviderIDAndInstanceName(provider.ID, instanceName) + instance, err := instanceDAO.GetByProviderIDAndInstanceName(ctx, db, provider.ID, instanceName) if err != nil { return nil, "", nil, 0, fmt.Errorf("instance %q lookup failed: %w", instanceName, err) } @@ -200,7 +200,7 @@ func resolveModelConfigFromProviderInstance(tenantID string, modelType entity.Mo baseURL := extra["base_url"] modelObj, modelErr := modelDAO.GetByProviderIDAndInstanceIDAndModelTypeAndModelName( - provider.ID, instance.ID, int(modelType), pureModelName, + ctx, db, provider.ID, instance.ID, int(modelType), pureModelName, ) switch { case modelErr == nil: diff --git a/internal/ingestion/component/extractor.go b/internal/ingestion/component/extractor.go index 588c272a41..f0d9ea28f3 100644 --- a/internal/ingestion/component/extractor.go +++ b/internal/ingestion/component/extractor.go @@ -540,7 +540,7 @@ func (c *ExtractorComponent) Invoke(ctx context.Context, db *gorm.DB, inputs map if err := runtime.WithTimeout(ctx, extractorTimeout, func(timeoutCtx context.Context) error { // Tag phase: run when auto_tags > 0 and we have chunks. if c.Param.AutoTags > 0 && len(in.chunks) > 0 { - tagged, tagErr := c.runAutoTags(timeoutCtx, in) + tagged, tagErr := c.runAutoTags(timeoutCtx, db, in) if tagErr != nil { return tagErr } @@ -548,7 +548,7 @@ func (c *ExtractorComponent) Invoke(ctx context.Context, db *gorm.DB, inputs map } if len(in.chunks) == 0 { - ans, callErr := c.call(timeoutCtx, in, "") + ans, callErr := c.call(timeoutCtx, db, in, "") if callErr != nil { return callErr } @@ -562,18 +562,18 @@ func (c *ExtractorComponent) Invoke(ctx context.Context, db *gorm.DB, inputs map } if c.Param.AutoKeywords > 0 { - if err := c.runAutoKeywords(timeoutCtx, in, ck, text); err != nil { + if err := c.runAutoKeywords(timeoutCtx, db, in, ck, text); err != nil { return fmt.Errorf("chunk %d keywords: %w", i, err) } } if c.Param.AutoQuestions > 0 { - if err := c.runAutoQuestions(timeoutCtx, in, ck, text); err != nil { + if err := c.runAutoQuestions(timeoutCtx, db, in, ck, text); err != nil { return fmt.Errorf("chunk %d questions: %w", i, err) } } if in.fieldName != "" { - ans, callErr := c.call(timeoutCtx, in, text) + ans, callErr := c.call(timeoutCtx, db, in, text) if callErr != nil { return fmt.Errorf("chunk %d: %w", i, callErr) } @@ -594,7 +594,7 @@ func (c *ExtractorComponent) Invoke(ctx context.Context, db *gorm.DB, inputs map }, nil } -func (c *ExtractorComponent) runAutoKeywords(ctx context.Context, in extractorInputs, ck map[string]any, chunkText string) error { +func (c *ExtractorComponent) runAutoKeywords(ctx context.Context, db *gorm.DB, in extractorInputs, ck map[string]any, chunkText string) error { if _, exists := ck["important_kwd"]; exists { return nil } @@ -602,7 +602,7 @@ func (c *ExtractorComponent) runAutoKeywords(ctx context.Context, in extractorIn kwIn.prompt = "Output: " kwIn.systemPrompt = fmt.Sprintf(autoKeywordPrompt, c.Param.AutoKeywords, chunkText) kwIn.fieldName = "" - result, err := c.call(ctx, kwIn, "") + result, err := c.call(ctx, db, kwIn, "") if err != nil { return err } @@ -624,7 +624,7 @@ func (c *ExtractorComponent) runAutoKeywords(ctx context.Context, in extractorIn return nil } -func (c *ExtractorComponent) runAutoQuestions(ctx context.Context, in extractorInputs, ck map[string]any, chunkText string) error { +func (c *ExtractorComponent) runAutoQuestions(ctx context.Context, db *gorm.DB, in extractorInputs, ck map[string]any, chunkText string) error { if _, exists := ck["question_kwd"]; exists { return nil } @@ -632,7 +632,7 @@ func (c *ExtractorComponent) runAutoQuestions(ctx context.Context, in extractorI qIn.prompt = "Output: " qIn.systemPrompt = fmt.Sprintf(autoQuestionPrompt, c.Param.AutoQuestions, chunkText) qIn.fieldName = "" - result, err := c.call(ctx, qIn, "") + result, err := c.call(ctx, db, qIn, "") if err != nil { return err } @@ -694,8 +694,8 @@ func splitKeywords(s string) []string { // (empty string in the no-chunk fast path). The result is the // raw string from the model — JSON parsing happens here so // callers can rely on a structured value downstream. -func (c *ExtractorComponent) call(ctx context.Context, in extractorInputs, chunkText string) (any, error) { - driver, modelName, apiKey, baseURL, err := resolveExtractorChatTarget(ctx, in.llmID) +func (c *ExtractorComponent) call(ctx context.Context, db *gorm.DB, in extractorInputs, chunkText string) (any, error) { + driver, modelName, apiKey, baseURL, err := resolveExtractorChatTarget(ctx, db, in.llmID) if err != nil { return nil, err } @@ -731,14 +731,14 @@ func (c *ExtractorComponent) call(ctx context.Context, in extractorInputs, chunk // api_key / base_url. The llm_id may be a bare tenant_model UUID or // a composite "model@provider" string. Errors from DAO resolution are // propagated so the caller sees the real failure reason. -func resolveExtractorChatTarget(ctx context.Context, llmID string) (driver, modelName, apiKey, baseURL string, err error) { +func resolveExtractorChatTarget(ctx context.Context, db *gorm.DB, llmID string) (driver, modelName, apiKey, baseURL string, err error) { if override := getExtractorChatTargetResolverOverride(); override != nil { if driver, modelName, apiKey, baseURL, ok := override(llmID); ok { return driver, modelName, apiKey, baseURL, nil } } - cfg, cfgErr := resolveExtractorChatConfig(ctx, llmID) + cfg, cfgErr := resolveExtractorChatConfig(ctx, db, llmID) if cfgErr != nil { return "", "", "", "", cfgErr } @@ -773,7 +773,7 @@ type extractorChatConfig struct { // // Returns nil error when there is no canvas state (unit tests) — // the caller's @ split fallback handles that case. -func resolveExtractorChatConfig(ctx context.Context, compositeLLMID string) (extractorChatConfig, error) { +func resolveExtractorChatConfig(ctx context.Context, db *gorm.DB, compositeLLMID string) (extractorChatConfig, error) { state, _, err := runtime.GetStateFromContext[*runtime.CanvasState](ctx) if err != nil || state == nil { return extractorChatConfig{}, nil @@ -792,13 +792,13 @@ func resolveExtractorChatConfig(ctx context.Context, compositeLLMID string) (ext // returns a clear error if the record doesn't exist. No need // for a separate pre-check — resolveModelConfig's redundant // GetByID dispatch check is also bypassed. - driver, modelName, apiConfig, _, err = resolveModelConfigByID(tid, entity.ModelTypeChat, compositeLLMID) + driver, modelName, apiConfig, _, err = resolveModelConfigByID(ctx, db, tid, entity.ModelTypeChat, compositeLLMID) if err != nil { return extractorChatConfig{}, fmt.Errorf("extractor: tenant model %q not found or not usable: %w", compositeLLMID, err) } } else { // Composite "model@provider" path: delegate to the shared dispatcher. - driver, modelName, apiConfig, _, err = resolveModelConfig(tid, entity.ModelTypeChat, compositeLLMID) + driver, modelName, apiConfig, _, err = resolveModelConfig(ctx, db, tid, entity.ModelTypeChat, compositeLLMID) if err != nil { return extractorChatConfig{}, fmt.Errorf("extractor: resolve model %q: %w", compositeLLMID, err) } diff --git a/internal/ingestion/component/extractor_tag.go b/internal/ingestion/component/extractor_tag.go index 73d692ffdd..f4c8a14988 100644 --- a/internal/ingestion/component/extractor_tag.go +++ b/internal/ingestion/component/extractor_tag.go @@ -16,6 +16,7 @@ import ( "time" "github.com/xuri/excelize/v2" + "gorm.io/gorm" "github.com/cespare/xxhash/v2" eschema "github.com/cloudwego/eino/schema" @@ -134,7 +135,7 @@ func (c *boundedTagCache) markRecentLocked(key string) { var tagSourceFileIndexCache = newBoundedTagCache(tagSourceCacheMax) -func (c *ExtractorComponent) runAutoTags(ctx context.Context, in extractorInputs) ([]map[string]any, error) { +func (c *ExtractorComponent) runAutoTags(ctx context.Context, db *gorm.DB, in extractorInputs) ([]map[string]any, error) { indexed, ok := c.resolveTagSource(ctx) if !ok || len(in.chunks) == 0 { common.Info("extractor tags: skipped", @@ -171,7 +172,7 @@ func (c *ExtractorComponent) runAutoTags(ctx context.Context, in extractorInputs } if len(docsToTag) > 0 && in.llmID != "" { - driver, model, apiKey, baseURL, err := resolveExtractorChatTarget(ctx, in.llmID) + driver, model, apiKey, baseURL, err := resolveExtractorChatTarget(ctx, db, in.llmID) if err != nil { common.Warn("extractor tag: resolve model failed, skipping LLM tagging", zap.Error(err)) } diff --git a/internal/ingestion/component/extractor_test.go b/internal/ingestion/component/extractor_test.go index 9394c9ae6d..85b4cc347a 100644 --- a/internal/ingestion/component/extractor_test.go +++ b/internal/ingestion/component/extractor_test.go @@ -19,6 +19,7 @@ package component import ( "context" "errors" + "ragflow/internal/dao" "strings" "sync" "sync/atomic" @@ -786,8 +787,9 @@ func TestIsBareTenantModelID(t *testing.T) { // TestResolveExtractorChatTarget_AtSplitFallback verifies the @ split // fallback path works without canvas state (unit test compatibility). func TestResolveExtractorChatTarget_AtSplitFallback(t *testing.T) { + ctx := t.Context() driver, modelName, apiKey, baseURL, err := resolveExtractorChatTarget( - context.Background(), "gpt-4o-mini@openai") + ctx, dao.DB, "gpt-4o-mini@openai") if err != nil { t.Fatalf("unexpected error: %v", err) } @@ -805,8 +807,9 @@ func TestResolveExtractorChatTarget_AtSplitFallback(t *testing.T) { // TestResolveExtractorChatTarget_NoDriver verifies a non-@ plain string // without canvas state returns no driver (passes through to Chat()). func TestResolveExtractorChatTarget_NoDriver(t *testing.T) { + ctx := t.Context() driver, modelName, _, _, err := resolveExtractorChatTarget( - context.Background(), "plain-name") + ctx, dao.DB, "plain-name") if err != nil { t.Fatalf("unexpected error: %v", err) } diff --git a/internal/ingestion/component/pdf_vision_dispatch.go b/internal/ingestion/component/pdf_vision_dispatch.go index b285e154f1..f98b6bb47c 100644 --- a/internal/ingestion/component/pdf_vision_dispatch.go +++ b/internal/ingestion/component/pdf_vision_dispatch.go @@ -488,7 +488,7 @@ func defaultPDFVisionModelResolver( driver, modelName, apiConfig, _, err := resolveTenantModelByType(ctx, db, tenantID, entity.ModelTypeImage2Text) return driver, modelName, apiConfig, err } - driver, modelName, apiConfig, _, err := resolveModelConfig(tenantID, entity.ModelTypeImage2Text, modelID) + driver, modelName, apiConfig, _, err := resolveModelConfig(ctx, db, tenantID, entity.ModelTypeImage2Text, modelID) return driver, modelName, apiConfig, err } diff --git a/internal/service/chat_pipeline.go b/internal/service/chat_pipeline.go index 5c311eebbe..2ac0c6bc83 100644 --- a/internal/service/chat_pipeline.go +++ b/internal/service/chat_pipeline.go @@ -1888,7 +1888,7 @@ func (s *ChatPipelineService) getLLMModelConfig(ctx context.Context, chat *entit // when the LLM is registered as such, otherwise CHAT. modelType := entity.ModelTypeChat modelTypeStr := "chat" - if modelTypes, mtErr := s.ModelProviderSvc.ResolveModelType(chat.TenantID, chat.LLMID); mtErr == nil { + if modelTypes, mtErr := s.ModelProviderSvc.ResolveModelType(ctx, chat.TenantID, chat.LLMID); mtErr == nil { for _, mt := range modelTypes { if mt == entity.ModelTypeImage2Text { modelType = entity.ModelTypeImage2Text diff --git a/internal/service/chunk/chunk.go b/internal/service/chunk/chunk.go index 8af33cbdd8..3caab50d5a 100644 --- a/internal/service/chunk/chunk.go +++ b/internal/service/chunk/chunk.go @@ -353,7 +353,7 @@ func (s *ChunkService) RetrievalTest(ctx context.Context, req *service.Retrieval embdID = kbRecords[0].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) + _, embdID, err = dao.LookupTenantLLMByName(ctx, dao.DB, 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) } diff --git a/internal/service/dataset/crud.go b/internal/service/dataset/crud.go index 56a9d0a436..c82c650ceb 100644 --- a/internal/service/dataset/crud.go +++ b/internal/service/dataset/crud.go @@ -395,7 +395,7 @@ func (d *DatasetService) ListDatasets(ctx context.Context, id, name string, page if tenantEmbdID == "" && isHexID(kb.EmbdID) { tenantEmbdID = kb.EmbdID } - item["embedding_model"] = service.ResolveTenantModelDisplayName(tenantEmbdID, kb.EmbdID, modelNameCache) + item["embedding_model"] = service.ResolveTenantModelDisplayName(ctx, dao.DB, tenantEmbdID, kb.EmbdID, modelNameCache) data = append(data, item) } diff --git a/internal/service/generator.go b/internal/service/generator.go index ae7a84df5e..d3e4b405a9 100644 --- a/internal/service/generator.go +++ b/internal/service/generator.go @@ -111,7 +111,7 @@ func CrossLanguages(ctx context.Context, tenantID string, llmID string, query st var err error if llmID != "" { - modelTypes, err := modelProviderSvc.ResolveModelType(tenantID, llmID) + modelTypes, err := modelProviderSvc.ResolveModelType(ctx, tenantID, llmID) if err != nil { return query, fmt.Errorf("failed to get model type: %w", err) } diff --git a/internal/service/llm.go b/internal/service/llm.go index 4b70dd4efa..d0bfe85d65 100644 --- a/internal/service/llm.go +++ b/internal/service/llm.go @@ -62,7 +62,7 @@ func (s *LLMService) GetMyLLMs(ctx context.Context, tenantID string, includeDeta result := make(map[string]MyLLMFactory) if includeDetails { - objs, err := s.tenantLLMDAO.ListAllByTenant(tenantID) + objs, err := s.tenantLLMDAO.ListAllByTenant(ctx, dao.DB, tenantID) if err != nil { return nil, err } @@ -108,7 +108,7 @@ func (s *LLMService) GetMyLLMs(ctx context.Context, tenantID string, includeDeta result[llmFactory] = factory } } else { - objs, err := s.tenantLLMDAO.GetMyLLMs(tenantID) + objs, err := s.tenantLLMDAO.GetMyLLMs(ctx, dao.DB, tenantID) if err != nil { return nil, err } @@ -171,7 +171,7 @@ func (s *LLMService) ListLLMs(ctx context.Context, tenantID string, modelType st "ModelScope": true, } - objs, err := s.tenantLLMDAO.ListAllByTenant(tenantID) + objs, err := s.tenantLLMDAO.ListAllByTenant(ctx, dao.DB, tenantID) if err != nil { return nil, err } @@ -374,14 +374,14 @@ func (s *LLMService) SetAPIKey(ctx context.Context, tenantID string, req *SetAPI } llmConfig["max_tokens"] = maxTokens - existingLLM, _ := s.tenantLLMDAO.GetByTenantFactoryAndModelName(tenantID, factory, llm.LLMName) + existingLLM, _ := s.tenantLLMDAO.GetByTenantFactoryAndModelName(ctx, dao.DB, tenantID, factory, llm.LLMName) if existingLLM != nil { updates := map[string]interface{}{ "api_key": req.APIKey, "api_base": baseURL, "max_tokens": maxTokens, } - dao.DB.Model(&entity.TenantLLM{}). + dao.DB.WithContext(ctx).Model(&entity.TenantLLM{}). Where("tenant_id = ? AND llm_factory = ? AND llm_name = ?", tenantID, factory, llm.LLMName). Updates(updates) } else { @@ -397,7 +397,7 @@ func (s *LLMService) SetAPIKey(ctx context.Context, tenantID string, req *SetAPI MaxTokens: maxTokens, Status: "1", } - s.tenantLLMDAO.Create(tenantLLM) + s.tenantLLMDAO.Create(ctx, dao.DB, tenantLLM) } } diff --git a/internal/service/memory.go b/internal/service/memory.go index 6ee0fbd7d3..7ae1c27276 100644 --- a/internal/service/memory.go +++ b/internal/service/memory.go @@ -33,6 +33,8 @@ import ( "ragflow/internal/engine" enginetypes "ragflow/internal/engine/types" "ragflow/internal/service/nlp" + + "gorm.io/gorm" ) const ( @@ -1469,8 +1471,8 @@ func (s *MemoryService) ListMemories(ctx context.Context, userID string, tenantI } memoryMap := map[string]interface{}{ "id": resp.ID, - "llm_id": ResolveTenantModelDisplayName(ptrStringValue(resp.TenantLLMID), resp.LLMID, modelNameCache), - "embd_id": ResolveTenantModelDisplayName(ptrStringValue(resp.TenantEmbdID), resp.EmbdID, modelNameCache), + "llm_id": ResolveTenantModelDisplayName(ctx, dao.DB, ptrStringValue(resp.TenantLLMID), resp.LLMID, modelNameCache), + "embd_id": ResolveTenantModelDisplayName(ctx, dao.DB, ptrStringValue(resp.TenantEmbdID), resp.EmbdID, modelNameCache), "name": resp.Name, "avatar": resp.Avatar, "tenant_id": resp.TenantID, @@ -1494,7 +1496,7 @@ func (s *MemoryService) ListMemories(ctx context.Context, userID string, tenantI // ResolveTenantModelDisplayName turns a tenant_model ID into // modelName@instance@provider. rawModelID is the API-facing fallback // stored on memory.llm_id / memory.embd_id / knowledgebase.embd_id. -func ResolveTenantModelDisplayName(tenantModelID, rawModelID string, cache map[string]string) string { +func ResolveTenantModelDisplayName(ctx context.Context, db *gorm.DB, tenantModelID, rawModelID string, cache map[string]string) string { tenantModelID = strings.TrimSpace(tenantModelID) rawModelID = strings.TrimSpace(rawModelID) if tenantModelID == "" || strings.Contains(tenantModelID, "@") { @@ -1510,15 +1512,15 @@ func ResolveTenantModelDisplayName(tenantModelID, rawModelID string, cache map[s cache[tenantModelID] = displayName }() - model, err := dao.NewTenantModelDAO().GetByID(tenantModelID) + model, err := dao.NewTenantModelDAO().GetByID(ctx, db, tenantModelID) if err != nil { return displayName } - instance, err := dao.NewTenantModelInstanceDAO().GetByID(model.InstanceID) + instance, err := dao.NewTenantModelInstanceDAO().GetByID(ctx, db, model.InstanceID) if err != nil { return displayName } - provider, err := dao.NewTenantModelProviderDAO().GetByID(model.ProviderID) + provider, err := dao.NewTenantModelProviderDAO().GetByID(ctx, db, model.ProviderID) if err != nil { return displayName } diff --git a/internal/service/model_service.go b/internal/service/model_service.go index e8c334c896..13a93b9014 100644 --- a/internal/service/model_service.go +++ b/internal/service/model_service.go @@ -149,7 +149,7 @@ type ModelProviderService struct { userTenantDAO *dao.UserTenantDAO } -// CheckConnectionRequest carries the credentials and optional instance selector +// CheckConnectionModelInfo CheckConnectionRequest carries the credentials and optional instance selector // for checking provider connectivity without creating a new model instance. type CheckConnectionModelInfo struct { ModelName string `json:"model_name"` @@ -166,7 +166,7 @@ type CheckConnectionRequest struct { ModelInfo []CheckConnectionModelInfo `json:"model_info"` } -func (m *ModelProviderService) AddModelProvider(providerName, userID string) (common.ErrorCode, error) { +func (m *ModelProviderService) AddModelProvider(ctx context.Context, providerName, userID string) (common.ErrorCode, error) { providerName = strings.TrimSpace(providerName) tenants, err := m.userTenantDAO.GetByUserIDAndRole(userID, "owner") @@ -180,7 +180,7 @@ func (m *ModelProviderService) AddModelProvider(providerName, userID string) (co tenantID := tenants[0].TenantID - existing, err := m.modelProviderDAO.GetByTenantIDAndProviderName(tenantID, providerName) + existing, err := m.modelProviderDAO.GetByTenantIDAndProviderName(ctx, dao.DB, tenantID, providerName) if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) { return common.CodeServerError, err } @@ -195,14 +195,14 @@ func (m *ModelProviderService) AddModelProvider(providerName, userID string) (co ProviderName: providerName, TenantID: tenantID, } - err = m.modelProviderDAO.Create(tenantModelProvider) + err = m.modelProviderDAO.Create(ctx, dao.DB, tenantModelProvider) if err != nil { return common.CodeServerError, fmt.Errorf("fail to create model provider: %s", err.Error()) } return common.CodeSuccess, nil } -func (m *ModelProviderService) ListProvidersOfTenant(userID string) ([]map[string]interface{}, common.ErrorCode, error) { +func (m *ModelProviderService) ListProvidersOfTenant(ctx context.Context, userID string) ([]map[string]interface{}, common.ErrorCode, error) { tenants, err := m.userTenantDAO.GetByUserIDAndRole(userID, "owner") if err != nil { @@ -215,7 +215,7 @@ func (m *ModelProviderService) ListProvidersOfTenant(userID string) ([]map[strin tenantID := tenants[0].TenantID - providerNames, err := m.modelProviderDAO.ListByID(tenantID) + providerNames, err := m.modelProviderDAO.ListByID(ctx, dao.DB, tenantID) if err != nil { return nil, common.CodeServerError, err } @@ -247,7 +247,7 @@ func (m *ModelProviderService) ListProvidersOfTenant(userID string) ([]map[strin // Set has_instance flag. Mirrors Python's: // provider_obj = TenantModelProviderService.get_by_tenant_id_and_provider_name(tenant_id, name) // has_instance = bool(provider_obj and TenantModelInstanceService.get_all_by_provider_id(provider_obj.id)) - provider["has_instance"] = m.providerHasInstance(tenantID, providerName) + provider["has_instance"] = m.providerHasInstance(ctx, tenantID, providerName) result = append(result, provider) } @@ -260,12 +260,12 @@ func (m *ModelProviderService) ListProvidersOfTenant(userID string) ([]map[strin // // provider_obj = TenantModelProviderService.get_by_tenant_id_and_provider_name(tenant_id, name) // has_instance = bool(provider_obj and TenantModelInstanceService.get_all_by_provider_id(provider_obj.id)) -func (m *ModelProviderService) providerHasInstance(tenantID, providerName string) bool { - providerObj, err := m.modelProviderDAO.GetByTenantIDAndProviderName(tenantID, providerName) +func (m *ModelProviderService) providerHasInstance(ctx context.Context, tenantID, providerName string) bool { + providerObj, err := m.modelProviderDAO.GetByTenantIDAndProviderName(ctx, dao.DB, tenantID, providerName) if err != nil { return false } - instances, err := m.modelInstanceDAO.GetAllInstancesByProviderID(providerObj.ID) + instances, err := m.modelInstanceDAO.GetAllInstancesByProviderID(ctx, dao.DB, providerObj.ID) if err != nil { return false } @@ -283,7 +283,7 @@ func isExcludedTenantProvider(name string) bool { return false } -func (m *ModelProviderService) DeleteModelProvider(userID, providerName string) (common.ErrorCode, error) { +func (m *ModelProviderService) DeleteModelProvider(ctx context.Context, userID, providerName string) (common.ErrorCode, error) { tenants, err := m.userTenantDAO.GetByUserIDAndRole(userID, "owner") if err != nil { return common.CodeServerError, err @@ -294,13 +294,13 @@ func (m *ModelProviderService) DeleteModelProvider(userID, providerName string) tenantID := tenants[0].TenantID // Find the provider first. - provider, err := m.modelProviderDAO.GetByTenantIDAndProviderName(tenantID, providerName) + provider, err := m.modelProviderDAO.GetByTenantIDAndProviderName(ctx, dao.DB, tenantID, providerName) if err != nil { return common.CodeNotFound, fmt.Errorf("provider %s not found", providerName) } // Delete all models and instances under this provider. - instances, err := m.modelInstanceDAO.GetAllInstancesByProviderID(provider.ID) + instances, err := m.modelInstanceDAO.GetAllInstancesByProviderID(ctx, dao.DB, provider.ID) if err != nil { return common.CodeServerError, err } @@ -309,15 +309,15 @@ func (m *ModelProviderService) DeleteModelProvider(userID, providerName string) for i, inst := range instances { instanceIDs[i] = inst.ID } - if _, err := m.modelDAO.DeleteByInstanceIDs(instanceIDs); err != nil { + if _, err = m.modelDAO.DeleteByInstanceIDs(ctx, dao.DB, instanceIDs); err != nil { return common.CodeServerError, err } - if _, err := m.modelInstanceDAO.DeleteByProviderID(provider.ID); err != nil { + if _, err = m.modelInstanceDAO.DeleteByProviderID(ctx, dao.DB, provider.ID); err != nil { return common.CodeServerError, err } } - _, err = m.modelProviderDAO.DeleteByTenantIDAndProviderName(tenantID, providerName) + _, err = m.modelProviderDAO.DeleteByTenantIDAndProviderName(ctx, dao.DB, tenantID, providerName) if err != nil { return common.CodeServerError, err } @@ -341,12 +341,12 @@ func (m *ModelProviderService) ListSupportedModels(ctx context.Context, provider tenantID := tenants[0].TenantID // Check if provider exists - provider, err := m.modelProviderDAO.GetByTenantIDAndProviderName(tenantID, providerName) + provider, err := m.modelProviderDAO.GetByTenantIDAndProviderName(ctx, dao.DB, tenantID, providerName) if err != nil { return nil, err } - instance, err := m.modelInstanceDAO.GetByProviderIDAndInstanceName(provider.ID, instanceName) + instance, err := m.modelInstanceDAO.GetByProviderIDAndInstanceName(ctx, dao.DB, provider.ID, instanceName) if err != nil { return nil, err } @@ -407,12 +407,12 @@ type CreateInstanceModelInfo struct { Extra map[string]interface{} `json:"extra"` } -func (m *ModelProviderService) getProviderByIDOrName(tenantID, providerIDOrName string) (*entity.TenantModelProvider, error) { - provider, err := m.modelProviderDAO.GetByID(providerIDOrName) +func (m *ModelProviderService) getProviderByIDOrName(ctx context.Context, tenantID, providerIDOrName string) (*entity.TenantModelProvider, error) { + provider, err := m.modelProviderDAO.GetByID(ctx, dao.DB, providerIDOrName) if err == nil && provider.TenantID == tenantID { return provider, nil } - return m.modelProviderDAO.GetByTenantIDAndProviderName(tenantID, strings.TrimSpace(providerIDOrName)) + return m.modelProviderDAO.GetByTenantIDAndProviderName(ctx, dao.DB, tenantID, strings.TrimSpace(providerIDOrName)) } func (m *ModelProviderService) CreateProviderInstance(ctx context.Context, providerIDOrName, instanceName, apiKey, baseURL, region, userID string, modelInfo []CreateInstanceModelInfo) (common.ErrorCode, error) { @@ -430,7 +430,7 @@ func (m *ModelProviderService) CreateProviderInstance(ctx context.Context, provi tenantID := tenants[0].TenantID - provider, err := m.getProviderByIDOrName(tenantID, providerIDOrName) + provider, err := m.getProviderByIDOrName(ctx, tenantID, providerIDOrName) if err != nil { return common.CodeNotFound, fmt.Errorf("provider '%s' does not exist", providerIDOrName) } @@ -465,7 +465,7 @@ func (m *ModelProviderService) CreateProviderInstance(ctx context.Context, provi Status: "active", Extra: extraStr, } - err = m.modelInstanceDAO.Create(tenantModelInstance) + err = m.modelInstanceDAO.Create(ctx, dao.DB, tenantModelInstance) if err != nil { return common.CodeServerError, fmt.Errorf("fail to create model instance: %s", err.Error()) } @@ -481,7 +481,7 @@ func (m *ModelProviderService) CreateProviderInstance(ctx context.Context, provi verifyStatus = entity.ModelVerifyUnknown } model.Extra["verify"] = verifyStatus - if err := m.addModelToInstance(tenantID, providerName, instanceName, model); err != nil { + if err = m.addModelToInstance(ctx, tenantID, providerName, instanceName, model); err != nil { return common.CodeServerError, err } } @@ -500,16 +500,16 @@ func (m *ModelProviderService) CreateProviderInstance(ctx context.Context, provi if verifyStatus == "" { verifyStatus = entity.ModelVerifyUnknown } - extra := map[string]interface{}{ + extraMap := map[string]interface{}{ "verify": verifyStatus, } if llm.Tools != nil { - extra["is_tools"] = llm.Tools.Support + extraMap["is_tools"] = llm.Tools.Support } if llm.Thinking != nil { - extra["thinking"] = llm.Thinking.DefaultValue + extraMap["thinking"] = llm.Thinking.DefaultValue } - if err := m.addModelToInstance(tenantID, providerName, instanceName, CreateInstanceModelInfo{ + if err = m.addModelToInstance(ctx, tenantID, providerName, instanceName, CreateInstanceModelInfo{ ModelName: llm.Name, ModelTypes: llm.ModelTypes, MaxTokens: func() int { @@ -518,7 +518,7 @@ func (m *ModelProviderService) CreateProviderInstance(ctx context.Context, provi } return 8192 }(), - Extra: extra, + Extra: extraMap, }); err != nil { return common.CodeServerError, err } @@ -531,7 +531,7 @@ func (m *ModelProviderService) CreateProviderInstance(ctx context.Context, provi // CreateNameOnlyProviderInstance creates a provider instance with only a name, // skipping API key validation and model creation. -func (m *ModelProviderService) CreateNameOnlyProviderInstance(providerIDOrName, instanceName, userID string) (common.ErrorCode, error) { +func (m *ModelProviderService) CreateNameOnlyProviderInstance(ctx context.Context, providerIDOrName, instanceName, userID string) (common.ErrorCode, error) { providerIDOrName = strings.TrimSpace(providerIDOrName) if instanceName == "default" { @@ -547,7 +547,7 @@ func (m *ModelProviderService) CreateNameOnlyProviderInstance(providerIDOrName, } tenantID := tenants[0].TenantID - provider, err := m.getProviderByIDOrName(tenantID, providerIDOrName) + provider, err := m.getProviderByIDOrName(ctx, tenantID, providerIDOrName) if err != nil { return common.CodeNotFound, fmt.Errorf("provider '%s' does not exist", providerIDOrName) } @@ -562,7 +562,7 @@ func (m *ModelProviderService) CreateNameOnlyProviderInstance(providerIDOrName, Status: "active", Extra: "{}", } - err = m.modelInstanceDAO.Create(tenantModelInstance) + err = m.modelInstanceDAO.Create(ctx, dao.DB, tenantModelInstance) if err != nil { return common.CodeServerError, fmt.Errorf("fail to create model instance: %s", err.Error()) } @@ -623,19 +623,19 @@ func (m *ModelProviderService) verifyProviderAPIKey(ctx context.Context, provide } // addModelToInstance creates a single model under the given provider instance. -func (m *ModelProviderService) addModelToInstance(tenantID, providerName, instanceName string, model CreateInstanceModelInfo) error { - provider, err := m.modelProviderDAO.GetByTenantIDAndProviderName(tenantID, providerName) +func (m *ModelProviderService) addModelToInstance(ctx context.Context, tenantID, providerName, instanceName string, model CreateInstanceModelInfo) error { + provider, err := m.modelProviderDAO.GetByTenantIDAndProviderName(ctx, dao.DB, tenantID, providerName) if err != nil { return fmt.Errorf("no provider found for provider '%s'", providerName) } - instance, err := m.modelInstanceDAO.GetByProviderIDAndInstanceName(provider.ID, instanceName) + instance, err := m.modelInstanceDAO.GetByProviderIDAndInstanceName(ctx, dao.DB, provider.ID, instanceName) if err != nil { return fmt.Errorf("no instance found for provider '%s' and instance '%s'", providerName, instanceName) } // Check for duplicate model. - _, err = m.modelDAO.GetModelByProviderIDAndInstanceIDAndModelName(provider.ID, instance.ID, model.ModelName) + _, err = m.modelDAO.GetModelByProviderIDAndInstanceIDAndModelName(ctx, dao.DB, provider.ID, instance.ID, model.ModelName) if err == nil { return fmt.Errorf("model '%s' already exists for provider '%s' and instance '%s'", model.ModelName, providerName, instanceName) } @@ -674,14 +674,14 @@ func (m *ModelProviderService) addModelToInstance(tenantID, providerName, instan Extra: string(extraBytes), } - if err := m.modelDAO.Create(tenantModel); err != nil { + if err = m.modelDAO.Create(ctx, dao.DB, tenantModel); err != nil { return fmt.Errorf("fail to create model '%s': %s", model.ModelName, err.Error()) } return nil } -func (m *ModelProviderService) ListProviderInstances(providerIDOrName, userID string) ([]map[string]interface{}, common.ErrorCode, error) { +func (m *ModelProviderService) ListProviderInstances(ctx context.Context, providerIDOrName, userID string) ([]map[string]interface{}, common.ErrorCode, error) { providerIDOrName = strings.TrimSpace(providerIDOrName) // Get tenant ID from user @@ -697,13 +697,13 @@ func (m *ModelProviderService) ListProviderInstances(providerIDOrName, userID st tenantID := tenants[0].TenantID // Check if provider exists — try by ID first, then by name. - provider, err := m.getProviderByIDOrName(tenantID, providerIDOrName) + provider, err := m.getProviderByIDOrName(ctx, tenantID, providerIDOrName) if err != nil { - return nil, common.CodeDataError, fmt.Errorf("No provider found for provider '%s'", providerIDOrName) + return nil, common.CodeDataError, fmt.Errorf("no provider found for provider '%s'", providerIDOrName) } // Check if provider exists - instances, err := m.modelInstanceDAO.GetAllInstancesByProviderID(provider.ID) + instances, err := m.modelInstanceDAO.GetAllInstancesByProviderID(ctx, dao.DB, provider.ID) if err != nil { return nil, common.CodeServerError, err } @@ -718,7 +718,7 @@ func (m *ModelProviderService) ListProviderInstances(providerIDOrName, userID st // Parse extra to extract region var extraFields map[string]string if instance.Extra != "" { - if err := json.Unmarshal([]byte(instance.Extra), &extraFields); err != nil { + if err = json.Unmarshal([]byte(instance.Extra), &extraFields); err != nil { return nil, common.CodeServerError, err } } @@ -741,7 +741,7 @@ func (m *ModelProviderService) ListProviderInstances(providerIDOrName, userID st return result, common.CodeSuccess, nil } -func (m *ModelProviderService) ShowProviderInstance(providerName, instanceIDOrName, userID string) (map[string]interface{}, common.ErrorCode, error) { +func (m *ModelProviderService) ShowProviderInstance(ctx context.Context, providerName, instanceIDOrName, userID string) (map[string]interface{}, common.ErrorCode, error) { providerName = strings.TrimSpace(providerName) providerName = strings.ToLower(providerName) @@ -761,9 +761,9 @@ func (m *ModelProviderService) ShowProviderInstance(providerName, instanceIDOrNa // provider_obj = TenantModelProviderService.get_by_tenant_id_and_provider_id(tenant_id, provider_id_or_name) // if not provider_obj: // provider_obj = TenantModelProviderService.get_by_tenant_id_and_provider_name(tenant_id, provider_id_or_name) - provider, err := m.modelProviderDAO.GetByTenantIDAndProviderName(tenantID, providerName) + provider, err := m.modelProviderDAO.GetByTenantIDAndProviderName(ctx, dao.DB, tenantID, providerName) if err != nil { - return nil, common.CodeDataError, fmt.Errorf("No provider found for provider '%s'", providerName) + return nil, common.CodeDataError, fmt.Errorf("no provider found for provider '%s'", providerName) } // Find the instance — try by ID first, then by name. @@ -773,12 +773,12 @@ func (m *ModelProviderService) ShowProviderInstance(providerName, instanceIDOrNa // if not instance_obj: // instance_obj = TenantModelInstanceService.get_by_provider_id_and_instance_name(provider_id, instance_id_or_name) var instance *entity.TenantModelInstance - instance, err = m.modelInstanceDAO.GetByID(instanceIDOrName) + instance, err = m.modelInstanceDAO.GetByID(ctx, dao.DB, instanceIDOrName) if err != nil || instance.ProviderID != provider.ID { - instance, err = m.modelInstanceDAO.GetByProviderIDAndInstanceName(provider.ID, instanceIDOrName) + instance, err = m.modelInstanceDAO.GetByProviderIDAndInstanceName(ctx, dao.DB, provider.ID, instanceIDOrName) } if err != nil { - return nil, common.CodeDataError, fmt.Errorf("No instance found for provider '%s' and instance '%s'", providerName, instanceIDOrName) + return nil, common.CodeDataError, fmt.Errorf("no instance found for provider '%s' and instance '%s'", providerName, instanceIDOrName) } // Parse extra fields. Mirrors Python's: @@ -821,12 +821,12 @@ func (m *ModelProviderService) ShowInstanceBalance(ctx context.Context, provider tenantID := tenants[0].TenantID // Check if provider exists - provider, err := m.modelProviderDAO.GetByTenantIDAndProviderName(tenantID, providerName) + provider, err := m.modelProviderDAO.GetByTenantIDAndProviderName(ctx, dao.DB, tenantID, providerName) if err != nil { return nil, common.CodeServerError, err } - instance, err := m.modelInstanceDAO.GetByProviderIDAndInstanceName(provider.ID, instanceName) + instance, err := m.modelInstanceDAO.GetByProviderIDAndInstanceName(ctx, dao.DB, provider.ID, instanceName) if err != nil { return nil, common.CodeServerError, err } @@ -901,7 +901,7 @@ func (m *ModelProviderService) CheckConnection(ctx context.Context, providerName // results to the database — mirrors Python's per-model update_model calls // inside the /connection/verify REST endpoint. if instanceID != "" && len(modelVerifyResult) > 0 { - if dbErr := m.updateModelVerifyResults(userID, providerName, instanceID, modelVerifyResult); dbErr != nil { + if dbErr := m.updateModelVerifyResults(ctx, userID, providerName, instanceID, modelVerifyResult); dbErr != nil { common.Logger.Error("failed to persist model verify results", zap.Error(dbErr)) } } @@ -917,7 +917,7 @@ func (m *ModelProviderService) CheckConnection(ctx context.Context, providerName // tenant_model table. It mirrors the Python update_model() called from the // /api/v1/providers//connection/verify endpoint when instance_id is // present in the request body. -func (m *ModelProviderService) updateModelVerifyResults(userID, providerName, instanceID string, modelVerifyResult map[string]string) error { +func (m *ModelProviderService) updateModelVerifyResults(ctx context.Context, userID, providerName, instanceID string, modelVerifyResult map[string]string) error { // Resolve tenant from user. userTenants, err := m.userTenantDAO.GetByUserID(userID) if err != nil || len(userTenants) == 0 { @@ -926,13 +926,14 @@ func (m *ModelProviderService) updateModelVerifyResults(userID, providerName, in tenantID := userTenants[0].TenantID // Resolve provider DB record from tenant + provider name. - provider, err := m.modelProviderDAO.GetByTenantIDAndProviderName(tenantID, providerName) + provider, err := m.modelProviderDAO.GetByTenantIDAndProviderName(ctx, dao.DB, tenantID, providerName) if err != nil { return fmt.Errorf("provider %s not found for tenant %s: %w", providerName, tenantID, err) } for modelName, verifyStatus := range modelVerifyResult { - modelObj, err := m.modelDAO.GetModelByProviderIDAndInstanceIDAndModelName(provider.ID, instanceID, modelName) + var modelObj *entity.TenantModel + modelObj, err = m.modelDAO.GetModelByProviderIDAndInstanceIDAndModelName(ctx, dao.DB, provider.ID, instanceID, modelName) if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { // No existing row — nothing to update (default is active). @@ -946,12 +947,13 @@ func (m *ModelProviderService) updateModelVerifyResults(userID, providerName, in _ = json.Unmarshal([]byte(modelObj.Extra), &extra) } extra["verify"] = verifyStatus - extraJSON, err := json.Marshal(extra) + var extraJSON []byte + extraJSON, err = json.Marshal(extra) if err != nil { return fmt.Errorf("failed to marshal extra for %s: %w", modelName, err) } - if err := m.modelDAO.UpdateByID(modelObj.ID, map[string]interface{}{ + if err = m.modelDAO.UpdateByID(ctx, dao.DB, modelObj.ID, map[string]interface{}{ "extra": string(extraJSON), }); err != nil { return fmt.Errorf("failed to update verify status for %s: %w", modelName, err) @@ -1169,7 +1171,7 @@ func verifyASRModel(ctx context.Context, driver modelModule.ModelDriver, modelNa tmpPath := tmpFile.Name() defer os.Remove(tmpPath) - if _, err := tmpFile.Write(wavData); err != nil { + if _, err = tmpFile.Write(wavData); err != nil { tmpFile.Close() return fmt.Errorf("failed to write test WAV: %w", err) } @@ -1230,12 +1232,12 @@ func (m *ModelProviderService) CheckInstanceConnection(ctx context.Context, prov tenantID := tenants[0].TenantID // Check if provider exists - provider, err := m.modelProviderDAO.GetByTenantIDAndProviderName(tenantID, providerName) + provider, err := m.modelProviderDAO.GetByTenantIDAndProviderName(ctx, dao.DB, tenantID, providerName) if err != nil { return common.CodeServerError, err } - instance, err := m.modelInstanceDAO.GetByProviderIDAndInstanceName(provider.ID, instanceName) + instance, err := m.modelInstanceDAO.GetByProviderIDAndInstanceName(ctx, dao.DB, provider.ID, instanceName) if err != nil { return common.CodeServerError, err } @@ -1290,12 +1292,12 @@ func (m *ModelProviderService) ListTasks(ctx context.Context, providerName, inst tenantID := tenants[0].TenantID // Check if provider exists - provider, err := m.modelProviderDAO.GetByTenantIDAndProviderName(tenantID, providerName) + provider, err := m.modelProviderDAO.GetByTenantIDAndProviderName(ctx, dao.DB, tenantID, providerName) if err != nil { return nil, common.CodeServerError, err } - instance, err := m.modelInstanceDAO.GetByProviderIDAndInstanceName(provider.ID, instanceName) + instance, err := m.modelInstanceDAO.GetByProviderIDAndInstanceName(ctx, dao.DB, provider.ID, instanceName) if err != nil { return nil, common.CodeServerError, err } @@ -1351,12 +1353,12 @@ func (m *ModelProviderService) ShowTask(ctx context.Context, providerName, insta tenantID := tenants[0].TenantID // Check if provider exists - provider, err := m.modelProviderDAO.GetByTenantIDAndProviderName(tenantID, providerName) + provider, err := m.modelProviderDAO.GetByTenantIDAndProviderName(ctx, dao.DB, tenantID, providerName) if err != nil { return nil, common.CodeServerError, err } - instance, err := m.modelInstanceDAO.GetByProviderIDAndInstanceName(provider.ID, instanceName) + instance, err := m.modelInstanceDAO.GetByProviderIDAndInstanceName(ctx, dao.DB, provider.ID, instanceName) if err != nil { return nil, common.CodeServerError, err } @@ -1440,16 +1442,16 @@ func (m *ModelProviderService) ListTenantAddedModels(ctx context.Context, userID } // Mirror Python's ensure_*_from_env calls. - _ = m.ensureMineruFromEnv(tenantID) - _ = m.ensurePaddleOCREnabledFromEnv(tenantID) - _ = m.ensureOpenDataLoaderFromEnv(tenantID) + _ = m.ensureMineruFromEnv(ctx, tenantID) + _ = m.ensurePaddleOCREnabledFromEnv(ctx, tenantID) + _ = m.ensureOpenDataLoaderFromEnv(ctx, tenantID) var modelTypeFilterBin entity.ModelType if modelTypeFilter != "" { modelTypeFilterBin = entity.ModelTypeFromString(modelTypeFilter) } - providers, err := m.modelProviderDAO.GetByTenantID(tenantID) + providers, err := m.modelProviderDAO.GetByTenantID(ctx, dao.DB, tenantID) if err != nil { return nil, common.CodeServerError, err } @@ -1464,7 +1466,7 @@ func (m *ModelProviderService) ListTenantAddedModels(ctx context.Context, userID providerInfoByID[p.ID] = p } - instances, err := m.modelInstanceDAO.GetByProviderIDs(providerIDs) + instances, err := m.modelInstanceDAO.GetByProviderIDs(ctx, dao.DB, providerIDs) if err != nil { return nil, common.CodeServerError, err } @@ -1480,7 +1482,7 @@ func (m *ModelProviderService) ListTenantAddedModels(ctx context.Context, userID } // Fetch tenant model records and filter by model_type if needed. - modelRecords, err := m.modelDAO.GetModelsByProviderIDsAndInstanceIDs(providerIDs, instanceIDs) + modelRecords, err := m.modelDAO.GetModelsByProviderIDsAndInstanceIDs(ctx, dao.DB, providerIDs, instanceIDs) if err != nil { return nil, common.CodeServerError, err } @@ -1509,7 +1511,7 @@ func (m *ModelProviderService) ListTenantAddedModels(ctx context.Context, userID } if combinedType != 0 { rec.ModelType = int(combinedType) - _ = m.modelDAO.UpdateByID(rec.ID, map[string]interface{}{"model_type": int(combinedType)}) + _ = m.modelDAO.UpdateByID(ctx, dao.DB, rec.ID, map[string]interface{}{"model_type": int(combinedType)}) } } @@ -1668,21 +1670,21 @@ func (m *ModelProviderService) resolveModelListTenant(ctx context.Context, userI // ensureMineruFromEnv mirrors Python's ensure_mineru_from_env. // It ensures a MinerU OCR provider instance exists when env vars are configured. -func (m *ModelProviderService) ensureMineruFromEnv(tenantID string) error { +func (m *ModelProviderService) ensureMineruFromEnv(ctx context.Context, tenantID string) error { config := collectEnvConfig(mineruEnvKeys, mineruDefaultConfig) - return m.ensureOCRProviderFromEnv(tenantID, "MinerU", "mineru-from-env", config) + return m.ensureOCRProviderFromEnv(ctx, tenantID, "MinerU", "mineru-from-env", config) } // ensurePaddleOCREnabledFromEnv mirrors Python's ensure_paddleocr_from_env. -func (m *ModelProviderService) ensurePaddleOCREnabledFromEnv(tenantID string) error { +func (m *ModelProviderService) ensurePaddleOCREnabledFromEnv(ctx context.Context, tenantID string) error { config := collectEnvConfig(paddleOCREnvKeys, paddleOCRDefaultConfig) - return m.ensureOCRProviderFromEnv(tenantID, "PaddleOCR", "paddleocr-from-env", config) + return m.ensureOCRProviderFromEnv(ctx, tenantID, "PaddleOCR", "paddleocr-from-env", config) } // ensureOpenDataLoaderFromEnv mirrors Python's ensure_opendataloader_from_env. -func (m *ModelProviderService) ensureOpenDataLoaderFromEnv(tenantID string) error { +func (m *ModelProviderService) ensureOpenDataLoaderFromEnv(ctx context.Context, tenantID string) error { config := collectEnvConfig(openDataLoaderEnvKeys, openDataLoaderDefaultConfig) - return m.ensureOCRProviderFromEnv(tenantID, "OpenDataLoader", "opendataloader-from-env", config) + return m.ensureOCRProviderFromEnv(ctx, tenantID, "OpenDataLoader", "opendataloader-from-env", config) } // env key / default config tables for the three OCR providers. @@ -1746,13 +1748,13 @@ func collectEnvConfig(envKeys []string, defaultConfig map[string]interface{}) ma // ensureOCRProviderFromEnv mirrors Python's _ensure_ocr_provider_from_env. // It finds or creates a provider, instance, and model for the given OCR provider. -func (m *ModelProviderService) ensureOCRProviderFromEnv(tenantID, providerName, modelName string, config map[string]interface{}) error { +func (m *ModelProviderService) ensureOCRProviderFromEnv(ctx context.Context, tenantID, providerName, modelName string, config map[string]interface{}) error { if config == nil { return nil } // 1. Find or create the provider. - provider, err := m.modelProviderDAO.GetByTenantIDAndProviderName(tenantID, providerName) + provider, err := m.modelProviderDAO.GetByTenantIDAndProviderName(ctx, dao.DB, tenantID, providerName) if err != nil { if !dao.IsNotFoundErr(err) { return fmt.Errorf("failed to get provider %s: %w", providerName, err) @@ -1763,7 +1765,7 @@ func (m *ModelProviderService) ensureOCRProviderFromEnv(tenantID, providerName, TenantID: tenantID, ProviderName: providerName, } - if err := m.modelProviderDAO.Create(provider); err != nil { + if err = m.modelProviderDAO.Create(ctx, dao.DB, provider); err != nil { return fmt.Errorf("failed to create provider %s: %w", providerName, err) } } @@ -1775,7 +1777,7 @@ func (m *ModelProviderService) ensureOCRProviderFromEnv(tenantID, providerName, } apiKey := string(apiKeyBytes) - instance, err := m.modelInstanceDAO.GetInstanceByApiKey(apiKey, provider.ID) + instance, err := m.modelInstanceDAO.GetInstanceByApiKey(ctx, dao.DB, apiKey, provider.ID) if err != nil { if !dao.IsNotFoundErr(err) { return fmt.Errorf("failed to get instance for %s: %w", providerName, err) @@ -1788,13 +1790,15 @@ func (m *ModelProviderService) ensureOCRProviderFromEnv(tenantID, providerName, APIKey: apiKey, Extra: "{}", } - if err := m.modelInstanceDAO.Create(instance); err != nil { + if err = m.modelInstanceDAO.Create(ctx, dao.DB, instance); err != nil { return fmt.Errorf("failed to create instance for %s: %w", providerName, err) } } // 3. Find or create the model. _, err = m.modelDAO.GetByProviderIDAndInstanceIDAndModelTypeAndModelName( + ctx, + dao.DB, provider.ID, instance.ID, int(entity.ModelTypeOCR), @@ -1804,7 +1808,8 @@ func (m *ModelProviderService) ensureOCRProviderFromEnv(tenantID, providerName, if !dao.IsNotFoundErr(err) { return fmt.Errorf("failed to get model for %s: %w", providerName, err) } - extraBytes, err := json.Marshal(map[string]int{"max_tokens": 0}) + var extraBytes []byte + extraBytes, err = json.Marshal(map[string]int{"max_tokens": 0}) if err != nil { return fmt.Errorf("failed to marshal extra for %s model: %w", providerName, err) } @@ -1818,7 +1823,7 @@ func (m *ModelProviderService) ensureOCRProviderFromEnv(tenantID, providerName, Status: "active", Extra: string(extraBytes), } - if err := m.modelDAO.Create(tenantModel); err != nil { + if err = m.modelDAO.Create(ctx, dao.DB, tenantModel); err != nil { return fmt.Errorf("failed to create model for %s: %w", providerName, err) } } @@ -1838,16 +1843,16 @@ func (m *ModelProviderService) AlterProviderInstance(ctx context.Context, userID } tenantID := tenants[0].TenantID - provider, err := m.getProviderByIDOrName(tenantID, providerIDOrName) + provider, err := m.getProviderByIDOrName(ctx, tenantID, providerIDOrName) if err != nil { return common.CodeNotFound, fmt.Errorf("provider '%s' does not exist", providerIDOrName) } providerName := provider.ProviderName // Find the instance — try by ID first, then by name. - instance, err := m.modelInstanceDAO.GetByID(instanceIDOrName) + instance, err := m.modelInstanceDAO.GetByID(ctx, dao.DB, instanceIDOrName) if err != nil || instance.ProviderID != provider.ID { - instance, err = m.modelInstanceDAO.GetByProviderIDAndInstanceName(provider.ID, instanceIDOrName) + instance, err = m.modelInstanceDAO.GetByProviderIDAndInstanceName(ctx, dao.DB, provider.ID, instanceIDOrName) } if err != nil { return common.CodeNotFound, fmt.Errorf("no instance found for provider '%s' and instance '%s'", providerName, instanceIDOrName) @@ -1882,7 +1887,7 @@ func (m *ModelProviderService) AlterProviderInstance(ctx context.Context, userID // Preserve existing extra fields not overwritten. existingExtra := make(map[string]interface{}) if instance.Extra != "" { - if err := json.Unmarshal([]byte(instance.Extra), &existingExtra); err != nil { + if err = json.Unmarshal([]byte(instance.Extra), &existingExtra); err != nil { return common.CodeServerError, err } } @@ -1894,7 +1899,7 @@ func (m *ModelProviderService) AlterProviderInstance(ctx context.Context, userID return common.CodeServerError, err } instanceUpdates["extra"] = string(extraBytes) - if err := m.modelInstanceDAO.UpdateByID(instance.ID, instanceUpdates); err != nil { + if err = m.modelInstanceDAO.UpdateByID(ctx, dao.DB, instance.ID, instanceUpdates); err != nil { return common.CodeServerError, fmt.Errorf("fail to update instance: %s", err.Error()) } @@ -1905,7 +1910,7 @@ func (m *ModelProviderService) AlterProviderInstance(ctx context.Context, userID } // Upsert models: add new ones, update existing ones, remove ones no longer selected. - existingModels, err := m.modelDAO.GetModelsByInstanceID(instance.ID) + existingModels, err := m.modelDAO.GetModelsByInstanceID(ctx, dao.DB, instance.ID) if err != nil { return common.CodeServerError, err } @@ -1930,7 +1935,7 @@ func (m *ModelProviderService) AlterProviderInstance(ctx context.Context, userID } } if len(idsToRemove) > 0 { - if _, err := m.modelDAO.DeleteByIDs(idsToRemove); err != nil { + if _, err = m.modelDAO.DeleteByIDs(ctx, dao.DB, idsToRemove); err != nil { return common.CodeServerError, err } } @@ -1977,16 +1982,16 @@ func (m *ModelProviderService) AlterProviderInstance(ctx context.Context, userID if mdl.MaxTokens > 0 { mergedExtra["max_tokens"] = mdl.MaxTokens } - extraBytes, _ := json.Marshal(mergedExtra) + extraBytes, _ = json.Marshal(mergedExtra) updates["extra"] = string(extraBytes) if len(updates) > 0 { - if err := m.modelDAO.UpdateByID(existingMdl.ID, updates); err != nil { + if err = m.modelDAO.UpdateByID(ctx, dao.DB, existingMdl.ID, updates); err != nil { return common.CodeServerError, err } } } else { // Add new model. - if err := m.addModelToInstance(tenantID, providerName, effectiveInstanceName, mdl); err != nil { + if err = m.addModelToInstance(ctx, tenantID, providerName, effectiveInstanceName, mdl); err != nil { return common.CodeServerError, err } } @@ -1996,7 +2001,7 @@ func (m *ModelProviderService) AlterProviderInstance(ctx context.Context, userID return common.CodeSuccess, nil } -func (m *ModelProviderService) DropProviderInstances(providerIDOrName, userID string, instanceIDOrNames []string) (common.ErrorCode, error) { +func (m *ModelProviderService) DropProviderInstances(ctx context.Context, providerIDOrName, userID string, instanceIDOrNames []string) (common.ErrorCode, error) { if len(instanceIDOrNames) == 0 { return common.CodeBadRequest, errors.New("instances is required") } @@ -2014,7 +2019,7 @@ func (m *ModelProviderService) DropProviderInstances(providerIDOrName, userID st tenantID := tenants[0].TenantID // Find provider by ID or name - provider, err := m.getProviderByIDOrName(tenantID, providerIDOrName) + provider, err := m.getProviderByIDOrName(ctx, tenantID, providerIDOrName) if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return common.CodeNotFound, fmt.Errorf("no provider found for provider %q", providerIDOrName) @@ -2031,7 +2036,7 @@ func (m *ModelProviderService) DropProviderInstances(providerIDOrName, userID st var instance *entity.TenantModelInstance // Try by ID first, then by name — same as Python. if idOrName != "" { - instance, err = m.modelInstanceDAO.GetByID(idOrName) + instance, err = m.modelInstanceDAO.GetByID(ctx, dao.DB, idOrName) if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) { return common.CodeServerError, err } @@ -2040,7 +2045,7 @@ func (m *ModelProviderService) DropProviderInstances(providerIDOrName, userID st } } if instance == nil { - instance, err = m.modelInstanceDAO.GetByProviderIDAndInstanceName(provider.ID, idOrName) + instance, err = m.modelInstanceDAO.GetByProviderIDAndInstanceName(ctx, dao.DB, provider.ID, idOrName) if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { notExistInstances = append(notExistInstances, idOrName) @@ -2059,17 +2064,17 @@ func (m *ModelProviderService) DropProviderInstances(providerIDOrName, userID st // Second pass: delete models and instances by IDs. // Mirrors Python's: delete_models_by_instance_ids(instance_ids) // TenantModelInstanceService.delete_by_ids(instance_ids) - if _, err := m.modelDAO.DeleteByInstanceIDs(instanceIDs); err != nil { + if _, err = m.modelDAO.DeleteByInstanceIDs(ctx, dao.DB, instanceIDs); err != nil { return common.CodeServerError, err } - if _, err := m.modelInstanceDAO.DeleteByIDs(instanceIDs); err != nil { + if _, err = m.modelInstanceDAO.DeleteByIDs(ctx, dao.DB, instanceIDs); err != nil { return common.CodeServerError, err } return common.CodeSuccess, nil } -func (m *ModelProviderService) DropInstanceModels(providerIDOrName, instanceIDOrName, userID string, modelNames []string) (common.ErrorCode, error) { +func (m *ModelProviderService) DropInstanceModels(ctx context.Context, providerIDOrName, instanceIDOrName, userID string, modelNames []string) (common.ErrorCode, error) { // Get tenant ID from user tenants, err := m.userTenantDAO.GetByUserIDAndRole(userID, "owner") if err != nil { @@ -2083,22 +2088,22 @@ func (m *ModelProviderService) DropInstanceModels(providerIDOrName, instanceIDOr tenantID := tenants[0].TenantID // Get provider by ID or name (matches Python's get_by_tenant_id_and_provider_id → get_by_tenant_id_and_provider_name fallback). - provider, err := m.getProviderByIDOrName(tenantID, providerIDOrName) + provider, err := m.getProviderByIDOrName(ctx, tenantID, providerIDOrName) if err != nil { - return common.CodeDataError, fmt.Errorf("No provider found for provider '%s'", providerIDOrName) + return common.CodeDataError, fmt.Errorf("no provider found for provider '%s'", providerIDOrName) } // Get instance by ID or name (matches Python's get_by_id → get_by_provider_id_and_instance_name fallback). - instance, err := m.modelInstanceDAO.GetByID(instanceIDOrName) + instance, err := m.modelInstanceDAO.GetByID(ctx, dao.DB, instanceIDOrName) if err != nil || instance.ProviderID != provider.ID { - instance, err = m.modelInstanceDAO.GetByProviderIDAndInstanceName(provider.ID, instanceIDOrName) + instance, err = m.modelInstanceDAO.GetByProviderIDAndInstanceName(ctx, dao.DB, provider.ID, instanceIDOrName) } if err != nil { - return common.CodeDataError, fmt.Errorf("No instance found for provider '%s' and instance '%s'", providerIDOrName, instanceIDOrName) + return common.CodeDataError, fmt.Errorf("no instance found for provider '%s' and instance '%s'", providerIDOrName, instanceIDOrName) } // Fetch all models under this instance and validate all requested names exist. - modelObjs, err := m.modelDAO.GetModelsByInstanceID(instance.ID) + modelObjs, err := m.modelDAO.GetModelsByInstanceID(ctx, dao.DB, instance.ID) if err != nil { return common.CodeServerError, err } @@ -2115,7 +2120,7 @@ func (m *ModelProviderService) DropInstanceModels(providerIDOrName, instanceIDOr } } if len(notExist) > 0 { - return common.CodeNotFound, fmt.Errorf("Models %v not found for provider '%s' and instance '%s'", notExist, providerIDOrName, instanceIDOrName) + return common.CodeNotFound, fmt.Errorf("models %v not found for provider '%s' and instance '%s'", notExist, providerIDOrName, instanceIDOrName) } // Collect IDs of only the requested models to delete. @@ -2130,14 +2135,14 @@ func (m *ModelProviderService) DropInstanceModels(providerIDOrName, instanceIDOr } } - if _, err := m.modelDAO.DeleteByIDs(idsToDelete); err != nil { + if _, err = m.modelDAO.DeleteByIDs(ctx, dao.DB, idsToDelete); err != nil { return common.CodeServerError, err } return common.CodeSuccess, nil } -func (m *ModelProviderService) ListInstanceModels(providerName, instanceName, userID string) ([]map[string]interface{}, error) { +func (m *ModelProviderService) ListInstanceModels(ctx context.Context, providerName, instanceName, userID string) ([]map[string]interface{}, error) { // Get tenant ID from user tenants, err := m.userTenantDAO.GetByUserIDAndRole(userID, "owner") if err != nil { @@ -2151,22 +2156,22 @@ func (m *ModelProviderService) ListInstanceModels(providerName, instanceName, us tenantID := tenants[0].TenantID // Find provider by ID or name - provider, err := m.getProviderByIDOrName(tenantID, providerName) + provider, err := m.getProviderByIDOrName(ctx, tenantID, providerName) if err != nil { return nil, err } // Find instance by ID first, then by name - instance, err := m.modelInstanceDAO.GetByID(instanceName) + instance, err := m.modelInstanceDAO.GetByID(ctx, dao.DB, instanceName) if err != nil || instance.ProviderID != provider.ID { - instance, err = m.modelInstanceDAO.GetByProviderIDAndInstanceName(provider.ID, instanceName) + instance, err = m.modelInstanceDAO.GetByProviderIDAndInstanceName(ctx, dao.DB, provider.ID, instanceName) if err != nil { return nil, err } } // Get all models for this instance - modelObjs, err := m.modelDAO.GetModelsByInstanceID(instance.ID) + modelObjs, err := m.modelDAO.GetModelsByInstanceID(ctx, dao.DB, instance.ID) if err != nil { return nil, err } @@ -2219,7 +2224,7 @@ func (m *ModelProviderService) ListInstanceModels(providerName, instanceName, us return modelList, nil } -func (m *ModelProviderService) AlterModel(providerName, instanceName, modelName, userID, modelID string, updateDict map[string]interface{}) (common.ErrorCode, error) { +func (m *ModelProviderService) AlterModel(ctx context.Context, providerName, instanceName, modelName, userID, modelID string, updateDict map[string]interface{}) (common.ErrorCode, error) { modelName = strings.TrimSpace(modelName) modelID = strings.TrimSpace(modelID) if modelName == "" && modelID == "" { @@ -2246,19 +2251,19 @@ func (m *ModelProviderService) AlterModel(providerName, instanceName, modelName, tenantID := tenants[0].TenantID // Check if provider exists - provider, err := m.getProviderByIDOrName(tenantID, providerName) + provider, err := m.getProviderByIDOrName(ctx, tenantID, providerName) if err != nil { return common.CodeServerError, err } - instance, err := m.modelInstanceDAO.GetByProviderIDAndInstanceName(provider.ID, instanceName) + instance, err := m.modelInstanceDAO.GetByProviderIDAndInstanceName(ctx, dao.DB, provider.ID, instanceName) if err != nil { return common.CodeServerError, err } var model *entity.TenantModel if modelName != "" { - model, err = m.modelDAO.GetModelByProviderIDAndInstanceIDAndModelName(provider.ID, instance.ID, modelName) + model, err = m.modelDAO.GetModelByProviderIDAndInstanceIDAndModelName(ctx, dao.DB, provider.ID, instance.ID, modelName) if err != nil { if !errors.Is(err, gorm.ErrRecordNotFound) { return common.CodeServerError, err @@ -2269,7 +2274,7 @@ func (m *ModelProviderService) AlterModel(providerName, instanceName, modelName, return common.CodeBadRequest, errors.New("model ID does not match model name") } } else { - model, err = m.modelDAO.GetByID(modelID) + model, err = m.modelDAO.GetByID(ctx, dao.DB, modelID) if err != nil { if !errors.Is(err, gorm.ErrRecordNotFound) { return common.CodeServerError, err @@ -2335,7 +2340,7 @@ func (m *ModelProviderService) AlterModel(providerName, instanceName, modelName, } if len(toUpdate) > 0 { - if err := m.modelDAO.UpdateByID(model.ID, toUpdate); err != nil { + if err = m.modelDAO.UpdateByID(ctx, dao.DB, model.ID, toUpdate); err != nil { return common.CodeServerError, err } } @@ -2445,7 +2450,7 @@ func maxTokensFromTenantModelExtra(modelEntity *entity.TenantModel, fallback int return fallback, nil } -func (m *ModelProviderService) getModelInstanceAndProviderByName(providerName, instanceName, modelName *string, userID string, apiConfig *modelModule.APIConfig) (*ModelInstanceAndProviderInfo, error) { +func (m *ModelProviderService) getModelInstanceAndProviderByName(ctx context.Context, providerName, instanceName, modelName *string, userID string, apiConfig *modelModule.APIConfig) (*ModelInstanceAndProviderInfo, error) { // Get tenant ID from user tenants, err := m.userTenantDAO.GetByUserIDAndRole(userID, "owner") if err != nil { @@ -2459,17 +2464,17 @@ func (m *ModelProviderService) getModelInstanceAndProviderByName(providerName, i tenantID := tenants[0].TenantID // Check if provider exists - providerEntity, err := m.modelProviderDAO.GetByTenantIDAndProviderName(tenantID, *providerName) + providerEntity, err := m.modelProviderDAO.GetByTenantIDAndProviderName(ctx, dao.DB, tenantID, *providerName) if err != nil { return nil, err } - instanceEntity, err := m.modelInstanceDAO.GetByProviderIDAndInstanceName(providerEntity.ID, *instanceName) + instanceEntity, err := m.modelInstanceDAO.GetByProviderIDAndInstanceName(ctx, dao.DB, providerEntity.ID, *instanceName) if err != nil { return nil, err } - modelEntity, err := m.modelDAO.GetModelByProviderIDAndInstanceIDAndModelName(providerEntity.ID, instanceEntity.ID, *modelName) + modelEntity, err := m.modelDAO.GetModelByProviderIDAndInstanceIDAndModelName(ctx, dao.DB, providerEntity.ID, instanceEntity.ID, *modelName) if err != nil { // Not found model modelEntity = nil @@ -2518,7 +2523,7 @@ func (m *ModelProviderService) getModelInstanceAndProviderByName(providerName, i return result, nil } -func (m *ModelProviderService) getModelInstanceAndProviderByID(modelID *string, userID string, apiConfig *modelModule.APIConfig) (*ModelInstanceAndProviderInfo, error) { +func (m *ModelProviderService) getModelInstanceAndProviderByID(ctx context.Context, modelID *string, userID string, apiConfig *modelModule.APIConfig) (*ModelInstanceAndProviderInfo, error) { // Get tenant ID from user tenants, err := m.userTenantDAO.GetByUserIDAndRole(userID, "owner") if err != nil { @@ -2531,17 +2536,17 @@ func (m *ModelProviderService) getModelInstanceAndProviderByID(modelID *string, tenantID := tenants[0].TenantID - modelEntity, err := m.modelDAO.GetByID(*modelID) + modelEntity, err := m.modelDAO.GetByID(ctx, dao.DB, *modelID) if err != nil { return nil, err } - instanceEntity, err := m.modelInstanceDAO.GetByID(modelEntity.InstanceID) + instanceEntity, err := m.modelInstanceDAO.GetByID(ctx, dao.DB, modelEntity.InstanceID) if err != nil { return nil, err } - providerEntity, err := m.modelProviderDAO.GetByID(instanceEntity.ProviderID) + providerEntity, err := m.modelProviderDAO.GetByID(ctx, dao.DB, instanceEntity.ProviderID) if err != nil { return nil, err } @@ -2600,12 +2605,12 @@ func (m *ModelProviderService) ChatToModelWithMessages(ctx context.Context, prov var info *ModelInstanceAndProviderInfo if modelID != nil { - info, err = m.getModelInstanceAndProviderByID(modelID, userID, apiConfig) + info, err = m.getModelInstanceAndProviderByID(ctx, modelID, userID, apiConfig) if err != nil || info == nil { return nil, common.CodeNotFound, err } } else { - info, err = m.getModelInstanceAndProviderByName(providerName, instanceName, modelName, userID, apiConfig) + info, err = m.getModelInstanceAndProviderByName(ctx, providerName, instanceName, modelName, userID, apiConfig) if err != nil || info == nil { return nil, common.CodeNotFound, err } @@ -2671,12 +2676,12 @@ func (m *ModelProviderService) ChatToModelStreamWithSender(ctx context.Context, var info *ModelInstanceAndProviderInfo if modelID != nil { - info, err = m.getModelInstanceAndProviderByID(modelID, userID, apiConfig) + info, err = m.getModelInstanceAndProviderByID(ctx, modelID, userID, apiConfig) if err != nil || info == nil { return common.CodeNotFound, err } } else { - info, err = m.getModelInstanceAndProviderByName(providerName, instanceName, modelName, userID, apiConfig) + info, err = m.getModelInstanceAndProviderByName(ctx, providerName, instanceName, modelName, userID, apiConfig) if err != nil || info == nil { return common.CodeNotFound, err } @@ -2764,12 +2769,12 @@ func (m *ModelProviderService) EmbedText(ctx context.Context, providerName, inst var info *ModelInstanceAndProviderInfo if modelID != nil { - info, err = m.getModelInstanceAndProviderByID(modelID, userID, apiConfig) + info, err = m.getModelInstanceAndProviderByID(ctx, modelID, userID, apiConfig) if err != nil || info == nil { return nil, common.CodeNotFound, err } } else { - info, err = m.getModelInstanceAndProviderByName(providerName, instanceName, modelName, userID, apiConfig) + info, err = m.getModelInstanceAndProviderByName(ctx, providerName, instanceName, modelName, userID, apiConfig) if err != nil || info == nil { return nil, common.CodeNotFound, err } @@ -2830,12 +2835,12 @@ func (m *ModelProviderService) RerankDocument(ctx context.Context, providerName, var info *ModelInstanceAndProviderInfo if modelID != nil { - info, err = m.getModelInstanceAndProviderByID(modelID, userID, apiConfig) + info, err = m.getModelInstanceAndProviderByID(ctx, modelID, userID, apiConfig) if err != nil || info == nil { return nil, common.CodeNotFound, err } } else { - info, err = m.getModelInstanceAndProviderByName(providerName, instanceName, modelName, userID, apiConfig) + info, err = m.getModelInstanceAndProviderByName(ctx, providerName, instanceName, modelName, userID, apiConfig) if err != nil || info == nil { return nil, common.CodeNotFound, err } @@ -2889,12 +2894,12 @@ func (m *ModelProviderService) TranscribeAudio(ctx context.Context, providerName var info *ModelInstanceAndProviderInfo if modelID != nil { - info, err = m.getModelInstanceAndProviderByID(modelID, userID, apiConfig) + info, err = m.getModelInstanceAndProviderByID(ctx, modelID, userID, apiConfig) if err != nil || info == nil { return nil, common.CodeNotFound, err } } else { - info, err = m.getModelInstanceAndProviderByName(providerName, instanceName, modelName, userID, apiConfig) + info, err = m.getModelInstanceAndProviderByName(ctx, providerName, instanceName, modelName, userID, apiConfig) if err != nil || info == nil { return nil, common.CodeNotFound, err } @@ -2946,12 +2951,12 @@ func (m *ModelProviderService) TranscribeAudioStream(ctx context.Context, provid var info *ModelInstanceAndProviderInfo if modelID != nil { - info, err = m.getModelInstanceAndProviderByID(modelID, userID, apiConfig) + info, err = m.getModelInstanceAndProviderByID(ctx, modelID, userID, apiConfig) if err != nil || info == nil { return common.CodeNotFound, err } } else { - info, err = m.getModelInstanceAndProviderByName(providerName, instanceName, modelName, userID, apiConfig) + info, err = m.getModelInstanceAndProviderByName(ctx, providerName, instanceName, modelName, userID, apiConfig) if err != nil || info == nil { return common.CodeNotFound, err } @@ -2999,12 +3004,12 @@ func (m *ModelProviderService) AudioSpeech(ctx context.Context, providerName, in var info *ModelInstanceAndProviderInfo if modelID != nil { - info, err = m.getModelInstanceAndProviderByID(modelID, userID, apiConfig) + info, err = m.getModelInstanceAndProviderByID(ctx, modelID, userID, apiConfig) if err != nil || info == nil { return nil, common.CodeNotFound, err } } else { - info, err = m.getModelInstanceAndProviderByName(providerName, instanceName, modelName, userID, apiConfig) + info, err = m.getModelInstanceAndProviderByName(ctx, providerName, instanceName, modelName, userID, apiConfig) if err != nil || info == nil { return nil, common.CodeNotFound, err } @@ -3055,12 +3060,12 @@ func (m *ModelProviderService) AudioSpeechStream(ctx context.Context, providerNa var info *ModelInstanceAndProviderInfo if modelID != nil { - info, err = m.getModelInstanceAndProviderByID(modelID, userID, apiConfig) + info, err = m.getModelInstanceAndProviderByID(ctx, modelID, userID, apiConfig) if err != nil || info == nil { return common.CodeNotFound, err } } else { - info, err = m.getModelInstanceAndProviderByName(providerName, instanceName, modelName, userID, apiConfig) + info, err = m.getModelInstanceAndProviderByName(ctx, providerName, instanceName, modelName, userID, apiConfig) if err != nil || info == nil { return common.CodeNotFound, err } @@ -3107,12 +3112,12 @@ func (m *ModelProviderService) OCRFile(ctx context.Context, providerName, instan var info *ModelInstanceAndProviderInfo if modelID != nil { - info, err = m.getModelInstanceAndProviderByID(modelID, userID, apiConfig) + info, err = m.getModelInstanceAndProviderByID(ctx, modelID, userID, apiConfig) if err != nil || info == nil { return nil, common.CodeNotFound, err } } else { - info, err = m.getModelInstanceAndProviderByName(providerName, instanceName, modelName, userID, apiConfig) + info, err = m.getModelInstanceAndProviderByName(ctx, providerName, instanceName, modelName, userID, apiConfig) if err != nil || info == nil { return nil, common.CodeNotFound, err } @@ -3163,12 +3168,12 @@ func (m *ModelProviderService) ParseFile(ctx context.Context, providerName, inst var info *ModelInstanceAndProviderInfo if modelID != nil { - info, err = m.getModelInstanceAndProviderByID(modelID, userID, apiConfig) + info, err = m.getModelInstanceAndProviderByID(ctx, modelID, userID, apiConfig) if err != nil || info == nil { return nil, common.CodeNotFound, err } } else { - info, err = m.getModelInstanceAndProviderByName(providerName, instanceName, modelName, userID, apiConfig) + info, err = m.getModelInstanceAndProviderByName(ctx, providerName, instanceName, modelName, userID, apiConfig) if err != nil || info == nil { return nil, common.CodeNotFound, err } @@ -3282,7 +3287,7 @@ func (m *ModelProviderService) GetModelConfigByID(ctx context.Context, userID st zap.String("modelType", modelType.String()), zap.String("modelID", modelID)) - modelEntity, err := m.modelDAO.GetByID(modelID) + modelEntity, err := m.modelDAO.GetByID(ctx, dao.DB, modelID) if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return nil, "", nil, 0, fmt.Errorf("tenant model id=%s not found", modelID) @@ -3296,7 +3301,7 @@ func (m *ModelProviderService) GetModelConfigByID(ctx context.Context, userID st return nil, "", nil, 0, fmt.Errorf("tenant model id=%s cannot be used as %s model", modelID, modelType.String()) } - providerEntity, err := m.modelProviderDAO.GetByID(modelEntity.ProviderID) + providerEntity, err := m.modelProviderDAO.GetByID(ctx, dao.DB, modelEntity.ProviderID) if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return nil, "", nil, 0, fmt.Errorf("provider id=%s not found for model id=%s", modelEntity.ProviderID, modelID) @@ -3324,7 +3329,7 @@ func (m *ModelProviderService) GetModelConfigByID(ctx context.Context, userID st } } - instanceEntity, err := m.modelInstanceDAO.GetByID(modelEntity.InstanceID) + instanceEntity, err := m.modelInstanceDAO.GetByID(ctx, dao.DB, modelEntity.InstanceID) if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return nil, "", nil, 0, fmt.Errorf("instance id=%s not found for model id=%s", modelEntity.InstanceID, modelID) @@ -3392,16 +3397,16 @@ func (m *ModelProviderService) ResolveModelConfig(ctx context.Context, tenantID if strings.TrimSpace(modelRef) == "" { return nil, "", nil, 0, fmt.Errorf("model ref is required") } - if _, err := m.modelDAO.GetByID(modelRef); err == nil { + if _, err := m.modelDAO.GetByID(ctx, dao.DB, modelRef); err == nil { 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) + return m.GetModelConfigFromProviderInstance(ctx, tenantID, modelType, modelRef) } 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, err := m.modelDAO.GetByID(ctx, dao.DB, modelName); err == nil { if modelObj.Status != "active" { return "", fmt.Errorf("tenant model id=%s is disabled", modelName) } @@ -3432,7 +3437,7 @@ func (m *ModelProviderService) ResolveModelID(ctx context.Context, tenantID stri return "", nil } - provider, err := m.modelProviderDAO.GetByTenantIDAndProviderName(tenantID, providerName) + provider, err := m.modelProviderDAO.GetByTenantIDAndProviderName(ctx, dao.DB, tenantID, providerName) if err != nil { return "", fmt.Errorf("provider %q lookup failed: %w", providerName, err) } @@ -3440,7 +3445,7 @@ func (m *ModelProviderService) ResolveModelID(ctx context.Context, tenantID stri return "", fmt.Errorf("provider %q not found for model %q", providerName, modelName) } - instance, err := m.modelInstanceDAO.GetByProviderIDAndInstanceName(provider.ID, instanceName) + instance, err := m.modelInstanceDAO.GetByProviderIDAndInstanceName(ctx, dao.DB, provider.ID, instanceName) if err != nil { return "", fmt.Errorf("instance %q lookup failed: %w", instanceName, err) } @@ -3448,7 +3453,7 @@ func (m *ModelProviderService) ResolveModelID(ctx context.Context, tenantID stri return "", fmt.Errorf("instance %q not found for model %q", instanceName, modelName) } - modelObj, err := m.modelDAO.GetByProviderIDAndInstanceIDAndModelTypeAndModelName(provider.ID, instance.ID, int(modelType), pureModelName) + modelObj, err := m.modelDAO.GetByProviderIDAndInstanceIDAndModelTypeAndModelName(ctx, dao.DB, provider.ID, instance.ID, int(modelType), pureModelName) if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return "", fmt.Errorf("model %q not found for model type %s", modelName, modelType.String()) @@ -3461,8 +3466,8 @@ func (m *ModelProviderService) ResolveModelID(ctx context.Context, tenantID stri return modelObj.ID, nil } -func (m *ModelProviderService) ResolveModelType(tenantID, modelRef string) ([]entity.ModelType, error) { - modelObj, err := m.modelDAO.GetByID(modelRef) +func (m *ModelProviderService) ResolveModelType(ctx context.Context, tenantID, modelRef string) ([]entity.ModelType, error) { + modelObj, err := m.modelDAO.GetByID(ctx, dao.DB, modelRef) if err == nil { if modelObj.Status != "active" { return nil, fmt.Errorf("tenant model id=%s is disabled", modelRef) @@ -3472,11 +3477,11 @@ func (m *ModelProviderService) ResolveModelType(tenantID, modelRef string) ([]en if !errors.Is(err, gorm.ErrRecordNotFound) { return nil, err } - return m.GetModelTypeByName(tenantID, modelRef) + return m.GetModelTypeByName(ctx, tenantID, modelRef) } // GetModelTypeByName returns the list of model types the given model is enrolled as. -func (m *ModelProviderService) GetModelTypeByName(tenantID, modelName string) ([]entity.ModelType, error) { +func (m *ModelProviderService) GetModelTypeByName(ctx context.Context, tenantID, modelName string) ([]entity.ModelType, error) { common.Debug("GetModelTypeByName", zap.String("tenantID", tenantID), zap.String("modelName", modelName)) @@ -3487,7 +3492,7 @@ func (m *ModelProviderService) GetModelTypeByName(tenantID, modelName string) ([ } // Direct provider lookup - provider, provErr := m.modelProviderDAO.GetByTenantIDAndProviderName(tenantID, providerName) + provider, provErr := m.modelProviderDAO.GetByTenantIDAndProviderName(ctx, dao.DB, tenantID, providerName) if provErr != nil { return nil, fmt.Errorf("provider %q lookup failed: %w", providerName, provErr) } @@ -3496,7 +3501,7 @@ func (m *ModelProviderService) GetModelTypeByName(tenantID, modelName string) ([ } // Direct instance lookup - instance, instErr := m.modelInstanceDAO.GetByProviderIDAndInstanceName(provider.ID, instanceName) + instance, instErr := m.modelInstanceDAO.GetByProviderIDAndInstanceName(ctx, dao.DB, provider.ID, instanceName) if instErr != nil { return nil, fmt.Errorf("instance %q lookup failed: %w", instanceName, instErr) } @@ -3505,7 +3510,7 @@ func (m *ModelProviderService) GetModelTypeByName(tenantID, modelName string) ([ } // Direct model lookup - modelObjs, modelErr := m.modelDAO.GetModelsByProviderIDAndInstanceIDAndModelName(provider.ID, instance.ID, pureModelName) + modelObjs, modelErr := m.modelDAO.GetModelsByProviderIDAndInstanceIDAndModelName(ctx, dao.DB, provider.ID, instance.ID, pureModelName) if modelErr == nil && len(modelObjs) > 0 { types := make([]entity.ModelType, 0, len(modelObjs)) for _, obj := range modelObjs { @@ -3561,7 +3566,7 @@ type ModelRequest struct { Thinking *bool `json:"thinking"` } -func (m *ModelProviderService) AddModel(request *AddModelRequest, userID string) (common.ErrorCode, error) { +func (m *ModelProviderService) AddModel(ctx context.Context, request *AddModelRequest, userID string) (common.ErrorCode, error) { if request == nil { return common.CodeBadRequest, errors.New("request is required") } @@ -3586,24 +3591,24 @@ func (m *ModelProviderService) AddModel(request *AddModelRequest, userID string) tenantID := tenants[0].TenantID // Get provider by ID or name (matches Python's get_by_tenant_id_and_provider_id → get_by_tenant_id_and_provider_name fallback). - provider, err := m.getProviderByIDOrName(tenantID, request.ProviderName) + provider, err := m.getProviderByIDOrName(ctx, tenantID, request.ProviderName) if err != nil { - return common.CodeDataError, fmt.Errorf("No provider found for provider '%s'", request.ProviderName) + return common.CodeDataError, fmt.Errorf("no provider found for provider '%s'", request.ProviderName) } // Get instance by ID or name (matches Python's get_by_id → get_by_provider_id_and_instance_name fallback). - instance, err := m.modelInstanceDAO.GetByID(request.InstanceName) + instance, err := m.modelInstanceDAO.GetByID(ctx, dao.DB, request.InstanceName) if err != nil || instance.ProviderID != provider.ID { - instance, err = m.modelInstanceDAO.GetByProviderIDAndInstanceName(provider.ID, request.InstanceName) + instance, err = m.modelInstanceDAO.GetByProviderIDAndInstanceName(ctx, dao.DB, provider.ID, request.InstanceName) } if err != nil { - return common.CodeDataError, fmt.Errorf("No instance found for provider '%s' and instance '%s'", request.ProviderName, request.InstanceName) + return common.CodeDataError, fmt.Errorf("no instance found for provider '%s' and instance '%s'", request.ProviderName, request.InstanceName) } // Check for duplicate model. - _, err = m.modelDAO.GetModelByProviderIDAndInstanceIDAndModelName(provider.ID, instance.ID, modelName) + _, err = m.modelDAO.GetModelByProviderIDAndInstanceIDAndModelName(ctx, dao.DB, provider.ID, instance.ID, modelName) if err == nil { - return common.CodeConflict, fmt.Errorf("Model '%s' already exists for provider '%s' and instance '%s'", modelName, request.ProviderName, request.InstanceName) + return common.CodeConflict, fmt.Errorf("model '%s' already exists for provider '%s' and instance '%s'", modelName, request.ProviderName, request.InstanceName) } if !errors.Is(err, gorm.ErrRecordNotFound) { return common.CodeServerError, err @@ -3669,7 +3674,7 @@ func (m *ModelProviderService) AddModel(request *AddModelRequest, userID string) Extra: string(extraBytes), } - if err := m.modelDAO.Create(tenantModel); err != nil { + if err = m.modelDAO.Create(ctx, dao.DB, tenantModel); err != nil { return common.CodeServerError, fmt.Errorf("fail to create model '%s': %s", modelName, err.Error()) } @@ -3683,7 +3688,7 @@ func (m *ModelProviderService) AddModel(request *AddModelRequest, userID string) // If the model is enrolled in tenant_model, that row is used (and INACTIVE rows // raise). Otherwise, the factory's LLM catalog is consulted, with // region=intl + siliconflow redirected to the siliconflow_intl factory. -func (m *ModelProviderService) GetModelConfigFromProviderInstance(tenantID string, modelType entity.ModelType, modelName string) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { +func (m *ModelProviderService) GetModelConfigFromProviderInstance(ctx context.Context, tenantID string, modelType entity.ModelType, modelName string) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { common.Debug("GetModelConfigFromProviderInstance", zap.String("tenantID", tenantID), zap.String("modelName", modelName), @@ -3770,7 +3775,7 @@ func (m *ModelProviderService) GetModelConfigFromProviderInstance(tenantID strin } // Direct provider lookup - provider, provErr := m.modelProviderDAO.GetByTenantIDAndProviderName(tenantID, providerName) + provider, provErr := m.modelProviderDAO.GetByTenantIDAndProviderName(ctx, dao.DB, tenantID, providerName) if provErr != nil { return nil, "", nil, 0, fmt.Errorf("provider %q lookup failed: %w", providerName, provErr) } @@ -3779,7 +3784,7 @@ func (m *ModelProviderService) GetModelConfigFromProviderInstance(tenantID strin } // Direct instance lookup - instance, instErr := m.modelInstanceDAO.GetByProviderIDAndInstanceName(provider.ID, instanceName) + instance, instErr := m.modelInstanceDAO.GetByProviderIDAndInstanceName(ctx, dao.DB, provider.ID, instanceName) if instErr != nil { return nil, "", nil, 0, fmt.Errorf("instance %q lookup failed: %w", instanceName, instErr) } @@ -3796,7 +3801,7 @@ func (m *ModelProviderService) GetModelConfigFromProviderInstance(tenantID strin // Direct model lookup modelObj, modelErr := m.modelDAO.GetByProviderIDAndInstanceIDAndModelTypeAndModelName( - provider.ID, instance.ID, int(modelType), pureModelName, + ctx, dao.DB, provider.ID, instance.ID, int(modelType), pureModelName, ) switch { case modelErr == nil: @@ -3881,7 +3886,7 @@ func (m *ModelProviderService) GetModelConfigFromProviderInstance(tenantID strin } // getModelConfig returns the model driver, model name, API config, and max tokens for a model -func (m *ModelProviderService) getModelConfig(tenantID, compositeModelName string) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { +func (m *ModelProviderService) getModelConfig(ctx context.Context, tenantID, compositeModelName string) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) { modelName, instanceName, providerName, err := parseModelName(compositeModelName) if err != nil { return nil, "", nil, 0, err @@ -3890,7 +3895,8 @@ func (m *ModelProviderService) getModelConfig(tenantID, compositeModelName strin // Check if provider exists (skip for Builtin provider) var providerID string if providerName != "Builtin" { - provider, err := m.modelProviderDAO.GetByTenantIDAndProviderName(tenantID, providerName) + var provider *entity.TenantModelProvider + provider, err = m.modelProviderDAO.GetByTenantIDAndProviderName(ctx, dao.DB, tenantID, providerName) if err != nil { return nil, "", nil, 0, err } @@ -3905,7 +3911,7 @@ func (m *ModelProviderService) getModelConfig(tenantID, compositeModelName strin // Get instance (skip for Builtin provider since it doesn't use tenant_model_instance) var instance *entity.TenantModelInstance if providerName != "Builtin" { - instance, err = m.modelInstanceDAO.GetByProviderIDAndInstanceName(providerID, instanceName) + instance, err = m.modelInstanceDAO.GetByProviderIDAndInstanceName(ctx, dao.DB, providerID, instanceName) if err != nil { return nil, "", nil, 0, err } @@ -3915,6 +3921,8 @@ func (m *ModelProviderService) getModelConfig(tenantID, compositeModelName strin common.Debug("getModelConfig instance found", zap.String("instanceName", instanceName)) } + // TODO: if provider name is Builtin, HOW TO? + var extra map[string]string var region string var baseURL string @@ -3960,7 +3968,7 @@ func (m *ModelProviderService) getModelConfig(tenantID, compositeModelName strin } var modelRecord *entity.TenantModel - modelRecord, err = m.modelDAO.GetModelByProviderIDAndInstanceIDAndModelName(providerID, instance.ID, modelName) + modelRecord, err = m.modelDAO.GetModelByProviderIDAndInstanceIDAndModelName(ctx, dao.DB, providerID, instance.ID, modelName) if err != nil { if !errors.Is(err, gorm.ErrRecordNotFound) { return nil, "", nil, 0, fmt.Errorf("tenant model %q lookup failed: %w", modelName, err) @@ -4006,11 +4014,11 @@ func (m *ModelProviderService) ShowModel(modelName string) (*modelModule.Model, // Returns false on lookup error or empty LLM ID so callers fall back to // chat — matches Python's branch order where only an EXPLICIT image2text // registration switches the model type away from chat. -func (m *ModelProviderService) isImage2TextLLM(tenantID, llmID string) bool { +func (m *ModelProviderService) isImage2TextLLM(ctx context.Context, tenantID, llmID string) bool { if m == nil || llmID == "" { return false } - modelTypes, err := m.ResolveModelType(tenantID, llmID) + modelTypes, err := m.ResolveModelType(ctx, tenantID, llmID) if err != nil { return false } @@ -4031,7 +4039,7 @@ func (m *ModelProviderService) GetChatModelConfig(ctx context.Context, tenantID return m.GetTenantDefaultModelByType(ctx, tenantID, entity.ModelTypeChat) } modelType := entity.ModelTypeChat - if m.isImage2TextLLM(tenantID, llmID) { + if m.isImage2TextLLM(ctx, tenantID, llmID) { modelType = entity.ModelTypeImage2Text } return m.ResolveModelConfig(ctx, tenantID, modelType, llmID) diff --git a/internal/service/model_service_test.go b/internal/service/model_service_test.go index 3c163ce547..6732480ee2 100644 --- a/internal/service/model_service_test.go +++ b/internal/service/model_service_test.go @@ -160,7 +160,8 @@ func TestModelProviderServiceAlterModelStatusByID(t *testing.T) { useModelProviderServiceTestDB(t, db) seedModelProviderServiceScope(t, db) - code, err := NewModelProviderService().AlterModel("OpenAI", "default", "", "user-1", "model-1", map[string]interface{}{"status": "inactive"}) + ctx := t.Context() + code, err := NewModelProviderService().AlterModel(ctx, "OpenAI", "default", "", "user-1", "model-1", map[string]interface{}{"status": "inactive"}) if err != nil { t.Fatalf("AlterModel() error = %v", err) } @@ -199,7 +200,8 @@ func TestModelProviderServiceGetModelConfigByID(t *testing.T) { } func TestModelProviderServiceAlterModelRejectsInvalidStatus(t *testing.T) { - code, err := NewModelProviderService().AlterModel("OpenAI", "default", "", "user-1", "model-1", map[string]interface{}{"status": "disabled"}) + ctx := t.Context() + code, err := NewModelProviderService().AlterModel(ctx, "OpenAI", "default", "", "user-1", "model-1", map[string]interface{}{"status": "disabled"}) if err == nil { t.Fatalf("AlterModel() error = nil, want invalid status error") } @@ -212,7 +214,8 @@ func TestModelProviderServiceAlterModelRejectsInvalidStatus(t *testing.T) { } func TestModelProviderServiceAlterModelRejectsMissingModelSelector(t *testing.T) { - code, err := NewModelProviderService().AlterModel("OpenAI", "default", "", "user-1", "", map[string]interface{}{"status": "active"}) + ctx := t.Context() + code, err := NewModelProviderService().AlterModel(ctx, "OpenAI", "default", "", "user-1", "", map[string]interface{}{"status": "active"}) if err == nil { t.Fatalf("AlterModel() error = nil, want missing model selector error") } @@ -232,7 +235,8 @@ func TestModelProviderServiceAlterModelRejectsWrongScopedModelID(t *testing.T) { t.Fatalf("failed to seed second instance: %v", err) } - code, err := NewModelProviderService().AlterModel("OpenAI", "other", "", "user-1", "model-1", map[string]interface{}{"status": "inactive"}) + ctx := t.Context() + code, err := NewModelProviderService().AlterModel(ctx, "OpenAI", "other", "", "user-1", "model-1", map[string]interface{}{"status": "inactive"}) if err == nil { t.Fatalf("AlterModel() error = nil, want not found error") } diff --git a/internal/service/openai_chat.go b/internal/service/openai_chat.go index 9d9347a2ac..23fcd34c67 100644 --- a/internal/service/openai_chat.go +++ b/internal/service/openai_chat.go @@ -236,7 +236,7 @@ func (s *OpenAIChatService) OpenAIChatCompletions(c *gin.Context, userID, chatID s.writeArgError(c, fmt.Sprintf("`llm_id` %s doesn't exist", req.Model)) return } - apiKey, apiErr := s.tenantLLMSvc.GetAPIKeyFromInstance(dialog.TenantID, req.Model) + apiKey, apiErr := s.tenantLLMSvc.GetAPIKeyFromInstance(ctx, dialog.TenantID, req.Model) if apiErr != nil || apiKey == "" { s.writeDataError(c, fmt.Sprintf("Cannot use specified model %s.", req.Model)) return diff --git a/internal/service/tenant.go b/internal/service/tenant.go index 84bf6fdc5d..78e03315b4 100644 --- a/internal/service/tenant.go +++ b/internal/service/tenant.go @@ -152,16 +152,16 @@ func NewTenantLLMService() *TenantLLMService { * // Get API key for model without factory * tenantLLM, err := service.GetAPIKey("tenant-123", "gpt-4") */ -func (s *TenantLLMService) GetAPIKey(tenantID, modelName string) (*entity.TenantLLM, error) { +func (s *TenantLLMService) GetAPIKey(ctx context.Context, tenantID, modelName string) (*entity.TenantLLM, error) { modelName, factory := s.SplitModelNameAndFactory(modelName) var tenantLLM *entity.TenantLLM var err error if factory == "" { - tenantLLM, err = s.tenantLLMDAO.GetByTenantIDAndLLMName(tenantID, modelName) + tenantLLM, err = s.tenantLLMDAO.GetByTenantIDAndLLMName(ctx, dao.DB, tenantID, modelName) } else { - tenantLLM, err = s.tenantLLMDAO.GetByTenantIDLLMNameAndFactory(tenantID, modelName, factory) + tenantLLM, err = s.tenantLLMDAO.GetByTenantIDLLMNameAndFactory(ctx, dao.DB, tenantID, modelName, factory) } if err != nil { @@ -186,7 +186,7 @@ func (s *TenantLLMService) SplitModelNameAndFactory(modelName string) (string, s // GetAPIKeyFromInstance returns the API key for the given composite model name // by looking it up in the tenant_model_instance table. compositeModelName is in // "model@instance@provider" or "model@provider" format. -func (s *TenantLLMService) GetAPIKeyFromInstance(tenantID, compositeModelName string) (string, error) { +func (s *TenantLLMService) GetAPIKeyFromInstance(ctx context.Context, tenantID, compositeModelName string) (string, error) { parts := strings.Split(compositeModelName, "@") if len(parts) < 2 { return "", fmt.Errorf("invalid model name format: %s", compositeModelName) @@ -204,7 +204,7 @@ func (s *TenantLLMService) GetAPIKeyFromInstance(tenantID, compositeModelName st return "", fmt.Errorf("invalid model name format: %s", compositeModelName) } - provider, err := s.modelProviderDAO.GetByTenantIDAndProviderName(tenantID, providerName) + provider, err := s.modelProviderDAO.GetByTenantIDAndProviderName(ctx, dao.DB, tenantID, providerName) if err != nil { return "", fmt.Errorf("provider %q not found: %w", providerName, err) } @@ -212,7 +212,7 @@ func (s *TenantLLMService) GetAPIKeyFromInstance(tenantID, compositeModelName st return "", fmt.Errorf("provider %q not found", providerName) } - instance, err := s.modelInstanceDAO.GetByProviderIDAndInstanceName(provider.ID, instanceName) + instance, err := s.modelInstanceDAO.GetByProviderIDAndInstanceName(ctx, dao.DB, provider.ID, instanceName) if err != nil { return "", fmt.Errorf("instance %q not found: %w", instanceName, err) } @@ -260,7 +260,7 @@ func (s *TenantLLMService) GetAPIKeyFromInstance(tenantID, compositeModelName st * // "tenant_embd_id": 456, // ID from tenant_llm table * // } */ -func (s *TenantLLMService) EnsureTenantModelIDForParams(tenantID string, params map[string]interface{}) map[string]interface{} { +func (s *TenantLLMService) EnsureTenantModelIDForParams(ctx context.Context, tenantID string, params map[string]interface{}) map[string]interface{} { paramKeys := []string{"llm_id", "embd_id", "asr_id", "img2txt_id", "rerank_id", "tts_id"} for _, key := range paramKeys { @@ -273,7 +273,7 @@ func (s *TenantLLMService) EnsureTenantModelIDForParams(tenantID string, params continue } - tenantLLM, err := s.GetAPIKey(tenantID, modelName) + tenantLLM, err := s.GetAPIKey(ctx, tenantID, modelName) if err == nil && tenantLLM != nil { params[tenantKey] = tenantLLM.ID } else { @@ -492,7 +492,7 @@ func (s *TenantService) GetDefaultModelName(ctx context.Context, tenantID string return modelID, nil } -func (s *TenantService) GetModelInfo(tenantID string, defaultModel string, modelType string) (*string, *string, *string, bool, error) { +func (s *TenantService) GetModelInfo(ctx context.Context, tenantID string, defaultModel string, modelType string) (*string, *string, *string, bool, error) { // Mirror Python's _get_model_info: right-anchored rsplit so that model // names containing '@' (e.g. LM Studio IDs like // "text-embedding-nomic-embed-text-v1.5@q8_0") remain intact. @@ -526,13 +526,13 @@ func (s *TenantService) GetModelInfo(tenantID string, defaultModel string, model } // Check if the provider exists for the tenant. - modelProvider, err := s.modelProviderDAO.GetByTenantIDAndProviderName(tenantID, providerName) + modelProvider, err := s.modelProviderDAO.GetByTenantIDAndProviderName(ctx, dao.DB, tenantID, providerName) if err != nil { return nil, nil, nil, false, err } // Check if the instance exists. - modelInstance, err := s.modelInstanceDAO.GetByProviderIDAndInstanceName(modelProvider.ID, instanceName) + modelInstance, err := s.modelInstanceDAO.GetByProviderIDAndInstanceName(ctx, dao.DB, modelProvider.ID, instanceName) if err != nil { return nil, nil, nil, false, err } @@ -548,7 +548,7 @@ func (s *TenantService) GetModelInfo(tenantID string, defaultModel string, model } // Check if the model exists and is active. - modelEntity, err := s.modelDAO.GetModelByProviderIDAndInstanceIDAndModelName(modelProvider.ID, modelInstance.ID, modelName) + modelEntity, err := s.modelDAO.GetModelByProviderIDAndInstanceIDAndModelName(ctx, dao.DB, modelProvider.ID, modelInstance.ID, modelName) if err != nil { if !dao.IsNotFoundErr(err) { return nil, nil, nil, false, err @@ -606,7 +606,7 @@ func (s *TenantService) ListTenantDefaultModels(ctx context.Context, userID stri var result []ModelItem - defaultChatModelProvider, defaultChatModelInstance, defaultChatModelName, defaultChatModelEnable, err := s.GetModelInfo(ownedTenant.TenantID, ownedTenant.LLMID, "chat") + defaultChatModelProvider, defaultChatModelInstance, defaultChatModelName, defaultChatModelEnable, err := s.GetModelInfo(ctx, ownedTenant.TenantID, ownedTenant.LLMID, "chat") if err == nil { result = append(result, ModelItem{ ModelProvider: defaultChatModelProvider, @@ -618,7 +618,7 @@ func (s *TenantService) ListTenantDefaultModels(ctx context.Context, userID stri }) } - defaultEmbeddingModelProvider, defaultEmbeddingModelInstance, defaultEmbeddingModelName, defaultEmbeddingModelEnable, err := s.GetModelInfo(ownedTenant.TenantID, ownedTenant.EmbDID, "embedding") + defaultEmbeddingModelProvider, defaultEmbeddingModelInstance, defaultEmbeddingModelName, defaultEmbeddingModelEnable, err := s.GetModelInfo(ctx, ownedTenant.TenantID, ownedTenant.EmbDID, "embedding") if err == nil { result = append(result, ModelItem{ ModelProvider: defaultEmbeddingModelProvider, @@ -630,7 +630,7 @@ func (s *TenantService) ListTenantDefaultModels(ctx context.Context, userID stri }) } - defaultRerankModelProvider, defaultRerankModelInstance, defaultRerankModelName, defaultRerankModelEnable, err := s.GetModelInfo(ownedTenant.TenantID, ownedTenant.RerankID, "rerank") + defaultRerankModelProvider, defaultRerankModelInstance, defaultRerankModelName, defaultRerankModelEnable, err := s.GetModelInfo(ctx, ownedTenant.TenantID, ownedTenant.RerankID, "rerank") if err == nil { result = append(result, ModelItem{ ModelProvider: defaultRerankModelProvider, @@ -642,7 +642,7 @@ func (s *TenantService) ListTenantDefaultModels(ctx context.Context, userID stri }) } - defaultASRModelProvider, defaultASRModelInstance, defaultASRModelName, defaultASREnable, err := s.GetModelInfo(ownedTenant.TenantID, ownedTenant.ASRID, "asr") + defaultASRModelProvider, defaultASRModelInstance, defaultASRModelName, defaultASREnable, err := s.GetModelInfo(ctx, ownedTenant.TenantID, ownedTenant.ASRID, "asr") if err == nil { result = append(result, ModelItem{ ModelProvider: defaultASRModelProvider, @@ -654,7 +654,7 @@ func (s *TenantService) ListTenantDefaultModels(ctx context.Context, userID stri }) } - defaultImage2TextModelProvider, defaultImage2TextModelInstance, defaultImage2TextModelName, defaultImage2TextModelEnable, err := s.GetModelInfo(ownedTenant.TenantID, ownedTenant.Img2TxtID, "vision") + defaultImage2TextModelProvider, defaultImage2TextModelInstance, defaultImage2TextModelName, defaultImage2TextModelEnable, err := s.GetModelInfo(ctx, ownedTenant.TenantID, ownedTenant.Img2TxtID, "vision") if err == nil { result = append(result, ModelItem{ ModelProvider: defaultImage2TextModelProvider, @@ -667,7 +667,7 @@ func (s *TenantService) ListTenantDefaultModels(ctx context.Context, userID stri } if ownedTenant.OCRID != "" { - defaultOCRModelProvider, defaultOCRModelInstance, defaultOCRModelName, defaultOCRModelEnable, err := s.GetModelInfo(ownedTenant.TenantID, ownedTenant.OCRID, "ocr") + defaultOCRModelProvider, defaultOCRModelInstance, defaultOCRModelName, defaultOCRModelEnable, err := s.GetModelInfo(ctx, ownedTenant.TenantID, ownedTenant.OCRID, "ocr") if err == nil { result = append(result, ModelItem{ ModelProvider: defaultOCRModelProvider, @@ -681,7 +681,7 @@ func (s *TenantService) ListTenantDefaultModels(ctx context.Context, userID stri } if ownedTenant.TTSID != "" { - defaultTTSModelProvider, defaultTTSModelInstance, defaultTTSModelName, defaultTTSModelEnable, err := s.GetModelInfo(ownedTenant.TenantID, ownedTenant.TTSID, "tts") + defaultTTSModelProvider, defaultTTSModelInstance, defaultTTSModelName, defaultTTSModelEnable, err := s.GetModelInfo(ctx, ownedTenant.TenantID, ownedTenant.TTSID, "tts") if err == nil { result = append(result, ModelItem{ ModelProvider: defaultTTSModelProvider, @@ -697,7 +697,7 @@ func (s *TenantService) ListTenantDefaultModels(ctx context.Context, userID stri return result, nil } -func (s *TenantService) checkModelAvailable(tenantID, providerName, instanceName, modelName, modelType string) error { +func (s *TenantService) checkModelAvailable(ctx context.Context, tenantID, providerName, instanceName, modelName, modelType string) error { _, _, modelTypeBit, err := tenantDefaultModelFields(modelType) if err != nil { return err @@ -709,7 +709,7 @@ func (s *TenantService) checkModelAvailable(tenantID, providerName, instanceName } // Static bypass: OCR with infiniflow@default@deepdoc is always enabled (mirrors Python _check_model_available). - if modelType == "ocr" && providerName == "infiniflow" && instanceName == "default" && modelName == "deepdoc" { + if modelType == "ocr" && providerName == "infiniflow" && instanceName == "default" { return nil } @@ -722,18 +722,18 @@ func (s *TenantService) checkModelAvailable(tenantID, providerName, instanceName } // Check if the provider and instance exists - modelProvider, err := s.modelProviderDAO.GetByTenantIDAndProviderName(tenantID, providerName) + modelProvider, err := s.modelProviderDAO.GetByTenantIDAndProviderName(ctx, dao.DB, tenantID, providerName) if err != nil { return err } - modelInstance, err := s.modelInstanceDAO.GetByProviderIDAndInstanceName(modelProvider.ID, instanceName) + modelInstance, err := s.modelInstanceDAO.GetByProviderIDAndInstanceName(ctx, dao.DB, modelProvider.ID, instanceName) if err != nil { return err } // Validate model availability through the DB (TenantModel table) - modelEntity, err := s.modelDAO.GetModelByProviderIDAndInstanceIDAndModelName(modelProvider.ID, modelInstance.ID, modelName) + modelEntity, err := s.modelDAO.GetModelByProviderIDAndInstanceIDAndModelName(ctx, dao.DB, modelProvider.ID, modelInstance.ID, modelName) if err != nil { if dao.IsNotFoundErr(err) { return fmt.Errorf("model %s isn't available", modelName) @@ -769,15 +769,15 @@ func (s *TenantService) SetTenantDefaultModels(ctx context.Context, userID, mode var tenantModelID interface{} if modelID != "" { - modelEntity, err := s.modelDAO.GetByID(modelID) + modelEntity, err := s.modelDAO.GetByID(ctx, dao.DB, modelID) if err != nil { return fmt.Errorf("model ID %s is invalid", modelID) } - instanceEntity, err := s.modelInstanceDAO.GetByID(modelEntity.InstanceID) + instanceEntity, err := s.modelInstanceDAO.GetByID(ctx, dao.DB, modelEntity.InstanceID) if err != nil { return fmt.Errorf("instance for model %s not found: %w", modelID, err) } - providerEntity, err := s.modelProviderDAO.GetByID(instanceEntity.ProviderID) + providerEntity, err := s.modelProviderDAO.GetByID(ctx, dao.DB, instanceEntity.ProviderID) if err != nil { return fmt.Errorf("provider for model %s not found: %w", modelID, err) } @@ -802,7 +802,7 @@ func (s *TenantService) SetTenantDefaultModels(ctx context.Context, userID, mode defaultModel = "" tenantModelID = nil } else if modelProvider != "" && modelInstance != "" && modelName != "" { - err = s.checkModelAvailable(ownedTenant.TenantID, modelProvider, modelInstance, modelName, modelType) + err = s.checkModelAvailable(ctx, ownedTenant.TenantID, modelProvider, modelInstance, modelName, modelType) if err != nil { return err } @@ -812,15 +812,15 @@ func (s *TenantService) SetTenantDefaultModels(ctx context.Context, userID, mode if modelProvider == "Builtin" { tenantModelID = nil } else { - modelProviderEntity, err := s.modelProviderDAO.GetByTenantIDAndProviderName(ownedTenant.TenantID, modelProvider) + modelProviderEntity, err := s.modelProviderDAO.GetByTenantIDAndProviderName(ctx, dao.DB, ownedTenant.TenantID, modelProvider) if err != nil { return err } - modelInstanceEntity, err := s.modelInstanceDAO.GetByProviderIDAndInstanceName(modelProviderEntity.ID, modelInstance) + modelInstanceEntity, err := s.modelInstanceDAO.GetByProviderIDAndInstanceName(ctx, dao.DB, modelProviderEntity.ID, modelInstance) if err != nil { return err } - modelEntity, err := s.modelDAO.GetModelByProviderIDAndInstanceIDAndModelName(modelProviderEntity.ID, modelInstanceEntity.ID, modelName) + modelEntity, err := s.modelDAO.GetModelByProviderIDAndInstanceIDAndModelName(ctx, dao.DB, modelProviderEntity.ID, modelInstanceEntity.ID, modelName) if err != nil { return err } diff --git a/internal/service/user.go b/internal/service/user.go index 99e6690ead..01376a30c5 100644 --- a/internal/service/user.go +++ b/internal/service/user.go @@ -893,7 +893,7 @@ func (s *UserService) SetTenantInfo(ctx context.Context, userID string, req *Set } tenantLLMService := NewTenantLLMService() - updates = tenantLLMService.EnsureTenantModelIDForParams(tenantID, updates) + updates = tenantLLMService.EnsureTenantModelIDForParams(ctx, tenantID, updates) if len(updates) > 0 { if err := tenantDAO.Update(ctx, dao.DB, tenantID, updates); err != nil {