mirror of
https://github.com/infiniflow/ragflow.git
synced 2026-07-23 17:06:42 +08:00
Fix: verify model api (#17183)
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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/<name>/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
|
||||
|
||||
Reference in New Issue
Block a user