From 340d30eb129b6640ffe0446e490963258df59a4e Mon Sep 17 00:00:00 2001 From: Lynn Date: Tue, 21 Jul 2026 19:09:46 +0800 Subject: [PATCH] Fix: verify model api (#17183) --- internal/handler/providers.go | 6 +- internal/service/model_service.go | 311 +++++++++++++++++++++++++++++- 2 files changed, 310 insertions(+), 7 deletions(-) diff --git a/internal/handler/providers.go b/internal/handler/providers.go index 47cb33c5ec..53eca76eed 100644 --- a/internal/handler/providers.go +++ b/internal/handler/providers.go @@ -301,7 +301,7 @@ func (h *ProviderHandler) CheckConnection(c *gin.Context) { } userID := c.GetString("user_id") - errCode, err := h.modelProviderService.CheckConnection(providerName, req.APIKey, req.Region, req.BaseURL, userID) + errCode, err := h.modelProviderService.CheckConnection(providerName, req.APIKey, req.Region, req.BaseURL, req.InstanceID, userID, service.ListModelNames(req.ModelInfo)) if err != nil { common.ErrorWithCode(c, errCode, err.Error()) return @@ -334,9 +334,9 @@ func (h *ProviderHandler) CheckInstanceConnection(c *gin.Context) { apikey, _ := instanceInfo["api_key"].(string) region, _ := instanceInfo["region"].(string) baseURL, _ := instanceInfo["base_url"].(string) + instanceID, _ := instanceInfo["id"].(string) - // Get tenant ID from user - errorCode, err := h.modelProviderService.CheckConnection(providerName, apikey, region, baseURL, userID) + errorCode, err := h.modelProviderService.CheckConnection(providerName, apikey, region, baseURL, instanceID, userID, nil) if err != nil { common.ErrorWithCode(c, errorCode, err.Error()) return diff --git a/internal/service/model_service.go b/internal/service/model_service.go index 302070eeaa..1eca432382 100644 --- a/internal/service/model_service.go +++ b/internal/service/model_service.go @@ -17,6 +17,7 @@ package service import ( + "encoding/binary" "encoding/json" "errors" "fmt" @@ -149,10 +150,29 @@ type ModelProviderService struct { // 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"` + ModelTypes []string `json:"model_type"` + MaxTokens int `json:"max_tokens"` + Extra map[string]interface{} `json:"extra"` +} + +func ListModelNames(modelInfo []CheckConnectionModelInfo) []string { + names := make([]string, 0, len(modelInfo)) + for _, mi := range modelInfo { + if mi.ModelName != "" { + names = append(names, mi.ModelName) + } + } + return names +} + type CheckConnectionRequest struct { - APIKey string `json:"api_key"` - Region string `json:"region"` - BaseURL string `json:"base_url"` + APIKey string `json:"api_key"` + Region string `json:"region"` + BaseURL string `json:"base_url"` + InstanceID string `json:"instance_id"` + ModelInfo []CheckConnectionModelInfo `json:"model_info"` } func (m *ModelProviderService) AddModelProvider(providerName, userID string) (common.ErrorCode, error) { @@ -850,7 +870,7 @@ func (m *ModelProviderService) ShowInstanceBalance(providerName, instanceName, u return result, common.CodeSuccess, nil } -func (m *ModelProviderService) CheckConnection(providerName, apiKey, region, baseURL string, userID string) (common.ErrorCode, error) { +func (m *ModelProviderService) CheckConnection(providerName, apiKey, region, baseURL, instanceID, userID string, modelInfo []string) (common.ErrorCode, error) { providerInfo := dao.GetModelProviderManager().FindProvider(providerName) if providerInfo == nil { return common.CodeServerError, fmt.Errorf("provider %s not found", providerName) @@ -886,9 +906,292 @@ func (m *ModelProviderService) CheckConnection(providerName, apiKey, region, bas return common.CodeServerError, err } + // Mirror Python verify_api_key: verify each model by making a real + // lightweight API request. Returns per-model verify results. + modelVerifyResult, verifyErr := verifyProviderModel(driver, providerInfo.Models, apiConfig, modelInfo) + + // When instanceID is provided (frontend passes it), persist the verify + // 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 { + common.Logger.Error("failed to persist model verify results", zap.Error(dbErr)) + } + } + + if verifyErr != nil { + return common.CodeServerError, verifyErr + } + return common.CodeSuccess, nil } +// updateModelVerifyResults persists the per-model verification status to the +// 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 { + // Resolve tenant from user. + userTenants, err := m.userTenantDAO.GetByUserID(userID) + if err != nil || len(userTenants) == 0 { + return fmt.Errorf("no tenant found for user %s", userID) + } + tenantID := userTenants[0].TenantID + + // Resolve provider DB record from tenant + provider name. + provider, err := m.modelProviderDAO.GetByTenantIDAndProviderName(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) + if err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + // No existing row — nothing to update (default is active). + continue + } + return fmt.Errorf("failed to look up %s for verify update: %w", modelName, err) + } + + extra := make(map[string]interface{}) + if modelObj.Extra != "" { + _ = json.Unmarshal([]byte(modelObj.Extra), &extra) + } + extra["verify"] = verifyStatus + 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{}{ + "extra": string(extraJSON), + }); err != nil { + return fmt.Errorf("failed to update verify status for %s: %w", modelName, err) + } + } + + return nil +} + +// verifyProviderModel mirrors Python verify_api_key's model-level verification. +// It tries each model registered for the provider in the factory JSON config +// and returns a map of modelName → verify status ("success"/"fail") so the +// caller can persist the results to the database. A nil error means at least +// one model passed verification. +func verifyProviderModel(driver modelModule.ModelDriver, providerModels []*modelModule.Model, apiConfig *modelModule.APIConfig, modelInfo []string) (map[string]string, error) { + modelVerifyResult := make(map[string]string) + + // Determine which models to verify: prefer the caller-supplied modelInfo + // list; fall back to the full provider model catalog. + var modelsToVerify []*modelModule.Model + if len(modelInfo) > 0 { + providerModelMap := make(map[string]*modelModule.Model, len(providerModels)) + for _, m := range providerModels { + providerModelMap[m.Name] = m + } + for _, name := range modelInfo { + name = strings.TrimSpace(name) + if m, ok := providerModelMap[name]; ok { + modelsToVerify = append(modelsToVerify, m) + } + } + } else { + modelsToVerify = providerModels + } + + if len(modelsToVerify) == 0 { + return modelVerifyResult, fmt.Errorf("no models found for provider") + } + + var errs []error + errSet := make(map[string]bool) + passedTypes := make(map[string]bool) + + for _, model := range modelsToVerify { + modelName := model.Name + anyPassed := false + + for _, modelType := range model.ModelTypes { + mtLower := strings.ToLower(modelType) + + // If a model type we've already verified successfully, skip. + if passedTypes[mtLower] { + continue + } + + var err error + + switch mtLower { + case "chat", "vision": + msg := []modelModule.Message{{Role: "user", Content: "Hi"}} + _, err = driver.ChatWithMessages(modelName, msg, apiConfig, nil, nil) + case "embedding": + _, err = driver.Embed(&modelName, []string{"test"}, apiConfig, nil, nil) + case "rerank": + _, err = driver.Rerank(&modelName, "test", []string{"test"}, apiConfig, &modelModule.RerankConfig{}, nil) + case "tts": + content := "hello" + _, err = driver.AudioSpeech(&modelName, &content, apiConfig, nil, nil) + case "asr": + err = verifyASRModel(driver, modelName, apiConfig) + case "ocr": + err = verifyOCRModel(driver, modelName, apiConfig) + default: + continue + } + + if err == nil { + passedTypes[mtLower] = true + anyPassed = true + break + } + + apiErr := extractAPIErrorMessage(err) + if !errSet[apiErr.Error()] { + errSet[apiErr.Error()] = true + errs = append(errs, apiErr) + } + } + + if anyPassed { + modelVerifyResult[modelName] = entity.ModelVerifySuccess + } else { + modelVerifyResult[modelName] = entity.ModelVerifyFail + } + } + + if len(passedTypes) == 0 { + return modelVerifyResult, fmt.Errorf("all model verification attempts failed: %w", errors.Join(errs...)) + } + + return modelVerifyResult, nil +} + +// extractAPIErrorMessage tries to parse the `message` field from a JSON error +// body embedded in a Go error string. If the body is valid JSON with a +// non-empty "message" key, the returned error contains only that message; +// otherwise the original error is returned unchanged. +func extractAPIErrorMessage(err error) error { + msg := err.Error() + // Look for the last '{'...'}' substring — that is typically the JSON body + // appended by API drivers like "API request failed with status 400: {...}". + start := strings.LastIndexByte(msg, '{') + if start < 0 { + return err + } + end := strings.LastIndexByte(msg, '}') + if end <= start { + return err + } + jsonStr := msg[start : end+1] + + var body struct { + Message string `json:"message"` + } + if json.Unmarshal([]byte(jsonStr), &body) != nil || body.Message == "" { + return err + } + return fmt.Errorf("%s", body.Message) +} + +// generateTestWAV creates a minimal silent WAV (16-bit mono PCM, 0.5 second, +// 16000 Hz sample rate) as a byte slice. Mirrors Python +// sequence2txt_model.py's _generate_test_wav: pure stdlib, no dependencies. +func generateTestWAV() []byte { + const ( + sampleRate = 16000 + durationSeconds = 0.5 + numChannels = 1 + bitsPerSample = 16 + ) + numSamples := int(sampleRate * durationSeconds) + dataSize := numSamples * numChannels * (bitsPerSample / 8) + + var buf []byte + + // RIFF header + buf = append(buf, []byte("RIFF")...) + buf = binary.LittleEndian.AppendUint32(buf, uint32(36+dataSize)) + buf = append(buf, []byte("WAVE")...) + + // fmt sub-chunk + buf = append(buf, []byte("fmt ")...) + buf = binary.LittleEndian.AppendUint32(buf, 16) // sub-chunk size + buf = binary.LittleEndian.AppendUint16(buf, 1) // PCM + buf = binary.LittleEndian.AppendUint16(buf, uint16(numChannels)) // mono + buf = binary.LittleEndian.AppendUint32(buf, uint32(sampleRate)) + buf = binary.LittleEndian.AppendUint32(buf, uint32(sampleRate*numChannels*bitsPerSample/8)) // byte rate + buf = binary.LittleEndian.AppendUint16(buf, uint16(numChannels*bitsPerSample/8)) // block align + buf = binary.LittleEndian.AppendUint16(buf, uint16(bitsPerSample)) + + // data sub-chunk + buf = append(buf, []byte("data")...) + buf = binary.LittleEndian.AppendUint32(buf, uint32(dataSize)) + buf = append(buf, make([]byte, dataSize)...) // silence + + return buf +} + +// verifyASRModel mirrors Python sequence2txt_model.py's check_available: +// generates a minimal test WAV, writes it to a temp file, calls +// TranscribeAudio, and checks the result for errors. +func verifyASRModel(driver modelModule.ModelDriver, modelName string, apiConfig *modelModule.APIConfig) error { + wavData := generateTestWAV() + + tmpFile, err := os.CreateTemp("", "ragflow-asr-verify-*.wav") + if err != nil { + return fmt.Errorf("failed to create temp WAV for ASR verification: %w", err) + } + tmpPath := tmpFile.Name() + defer os.Remove(tmpPath) + + if _, err := tmpFile.Write(wavData); err != nil { + tmpFile.Close() + return fmt.Errorf("failed to write test WAV: %w", err) + } + tmpFile.Close() + + resp, err := driver.TranscribeAudio(&modelName, &tmpPath, apiConfig, nil, nil) + if err != nil { + return err + } + if resp == nil || resp.Text == "" { + return fmt.Errorf("ASR model %s returned no transcription", modelName) + } + return nil +} + +// verifyOCRModel mirrors Python OCRModel.check_available by sending a +// minimal PNG (1×1 pixel white) through the OCR pipeline. Most OCR +// providers return an error or empty text for such a trivial image; +// we accept any non-error response as a successful connectivity check. +func verifyOCRModel(driver modelModule.ModelDriver, modelName string, apiConfig *modelModule.APIConfig) error { + // Send a minimal 1×1 white PNG through the OCR pipeline to verify + // connectivity. Most OCRModel drivers only check server reachability + // rather than performing full document parsing. + _, err := driver.OCRFile(&modelName, minimalPNG(), nil, apiConfig, nil, nil) + if err != nil { + return err + } + return nil +} + +// minimalPNG returns a 1×1 white PNG as a byte slice for OCR verification. +func minimalPNG() []byte { + return []byte{ + 0x89, 0x50, 0x4E, 0x47, 0x0D, 0x0A, 0x1A, 0x0A, + 0x00, 0x00, 0x00, 0x0D, 0x49, 0x48, 0x44, 0x52, + 0x00, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x01, + 0x08, 0x02, 0x00, 0x00, 0x00, 0x90, 0x77, 0x53, + 0xDE, 0x00, 0x00, 0x00, 0x0C, 0x49, 0x44, 0x41, + 0x54, 0x08, 0xD7, 0x63, 0x60, 0x60, 0xF8, 0x0F, + 0x00, 0x01, 0x01, 0x00, 0x05, 0x18, 0xD8, 0x32, + 0x48, 0x00, 0x00, 0x00, 0x00, 0x49, 0x45, 0x4E, + 0x44, 0xAE, 0x42, 0x60, 0x82, + } +} + func (m *ModelProviderService) CheckInstanceConnection(providerName, instanceName, userID string) (common.ErrorCode, error) { // Get tenant ID from user