mirror of
https://github.com/infiniflow/ragflow.git
synced 2026-08-05 15:20:30 +08:00
fix: XunFei Spark provider API key verification and model listing (#17652)
This commit is contained in:
@@ -17,6 +17,8 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
@@ -242,12 +244,32 @@ func (h *ProviderHandler) ShowModel(c *gin.Context) {
|
||||
|
||||
type CreateProviderInstanceRequest struct {
|
||||
InstanceName string `json:"instance_name" binding:"required"`
|
||||
APIKey string `json:"api_key"`
|
||||
APIKey json.RawMessage `json:"api_key"`
|
||||
BaseURL string `json:"base_url"`
|
||||
Region string `json:"region"`
|
||||
ModelInfo []service.CreateInstanceModelInfo `json:"model_info"`
|
||||
}
|
||||
|
||||
// normalizeAPIKey accepts api_key as either a JSON string or a JSON object
|
||||
// (credential bundles such as XunFei Spark's
|
||||
// {"spark_api_password": ..., "spark_app_id": ..., ...}) and normalizes it to
|
||||
// the string form persisted on the instance.
|
||||
func normalizeAPIKey(raw json.RawMessage) string {
|
||||
trimmed := strings.TrimSpace(string(raw))
|
||||
if trimmed == "" || trimmed == "null" {
|
||||
return ""
|
||||
}
|
||||
var s string
|
||||
if err := json.Unmarshal(raw, &s); err == nil {
|
||||
return s
|
||||
}
|
||||
var buf bytes.Buffer
|
||||
if err := json.Compact(&buf, raw); err != nil {
|
||||
return trimmed
|
||||
}
|
||||
return buf.String()
|
||||
}
|
||||
|
||||
func (h *ProviderHandler) CreateProviderInstance(c *gin.Context) {
|
||||
providerName := c.Param("provider_id_or_name")
|
||||
if providerName == "" {
|
||||
@@ -263,11 +285,12 @@ func (h *ProviderHandler) CreateProviderInstance(c *gin.Context) {
|
||||
}
|
||||
|
||||
userID := c.GetString("user_id")
|
||||
apiKey := normalizeAPIKey(req.APIKey)
|
||||
|
||||
// If the request body only contains "instance_name", create a name-only
|
||||
// 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 {
|
||||
if apiKey == "" && req.BaseURL == "" && req.Region == "" && len(req.ModelInfo) == 0 {
|
||||
code, err := h.modelProviderService.CreateNameOnlyProviderInstance(ctx, providerName, req.InstanceName, userID)
|
||||
if err != nil {
|
||||
common.ErrorWithCode(c, code, err.Error())
|
||||
@@ -277,7 +300,7 @@ func (h *ProviderHandler) CreateProviderInstance(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
_, err := h.modelProviderService.CreateProviderInstance(ctx, providerName, req.InstanceName, req.APIKey, req.BaseURL, req.Region, userID, req.ModelInfo)
|
||||
_, err := h.modelProviderService.CreateProviderInstance(ctx, providerName, req.InstanceName, apiKey, req.BaseURL, req.Region, userID, req.ModelInfo)
|
||||
if err != nil {
|
||||
common.ErrorWithCode(c, common.CodeServerError, err.Error())
|
||||
return
|
||||
@@ -371,7 +394,7 @@ func (h *ProviderHandler) CheckConnection(c *gin.Context) {
|
||||
}
|
||||
|
||||
userID := c.GetString("user_id")
|
||||
errCode, err := h.modelProviderService.CheckConnection(ctx, providerName, req.APIKey, req.Region, req.BaseURL, req.InstanceID, userID, req.ModelInfo)
|
||||
errCode, err := h.modelProviderService.CheckConnection(ctx, providerName, normalizeAPIKey(req.APIKey), req.Region, req.BaseURL, req.InstanceID, userID, req.ModelInfo)
|
||||
if err != nil {
|
||||
common.ErrorWithCode(c, errCode, err.Error())
|
||||
return
|
||||
@@ -475,7 +498,7 @@ func (h *ProviderHandler) ShowTask(c *gin.Context) {
|
||||
|
||||
type AlterProviderInstanceRequest struct {
|
||||
InstanceName string `json:"instance_name"`
|
||||
APIKey string `json:"api_key"`
|
||||
APIKey json.RawMessage `json:"api_key"`
|
||||
BaseURL string `json:"base_url"`
|
||||
Region string `json:"region"`
|
||||
ModelInfo []service.CreateInstanceModelInfo `json:"model_info"`
|
||||
@@ -513,7 +536,7 @@ func (h *ProviderHandler) AlterProviderInstance(c *gin.Context) {
|
||||
verify = *req.Verify
|
||||
}
|
||||
|
||||
code, err := h.modelProviderService.AlterProviderInstance(ctx, userID, providerName, instanceName, req.InstanceName, req.APIKey, req.BaseURL, req.Region, req.ModelInfo, verify)
|
||||
code, err := h.modelProviderService.AlterProviderInstance(ctx, userID, providerName, instanceName, req.InstanceName, normalizeAPIKey(req.APIKey), req.BaseURL, req.Region, req.ModelInfo, verify)
|
||||
if err != nil {
|
||||
common.ErrorWithCode(c, code, err.Error())
|
||||
return
|
||||
|
||||
Reference in New Issue
Block a user