package controllers import ( "encoding/json" "fmt" "io" "strconv" "strings" "sync" "time" "server/models" "server/pkg/jwtutil" beego "github.com/beego/beego/v2/server/web" ) // BackendAiProviderController AI接入配置控制器 type BackendAiProviderController struct { beego.Controller } func (c *BackendAiProviderController) aiClaims() (*jwtutil.Claims, error) { auth := c.Ctx.Request.Header.Get("Authorization") if auth == "" { return nil, fmt.Errorf("未登录") } parts := strings.SplitN(auth, " ", 2) if len(parts) != 2 || parts[0] != "Bearer" { return nil, fmt.Errorf("认证信息格式错误") } claims, err := jwtutil.ParseToken(parts[1]) if err != nil { return nil, fmt.Errorf("无效的token") } if claims.UserType != "backend" { return nil, fmt.Errorf("无权访问") } return claims, nil } func (c *BackendAiProviderController) aiJsonErr(httpStatus, bizCode int, msg string) { c.Ctx.Output.SetStatus(httpStatus) c.Data["json"] = map[string]interface{}{"code": bizCode, "msg": msg} _ = c.ServeJSON() } func (c *BackendAiProviderController) aiOk(data interface{}) { c.Data["json"] = map[string]interface{}{"code": 200, "msg": "success", "data": data} _ = c.ServeJSON() } type aiProviderPayload struct { ProviderType string `json:"provider_type"` Name string `json:"name"` ApiBase string `json:"api_base"` ApiKey string `json:"api_key"` Models []string `json:"models"` IsDefault int8 `json:"is_default"` Status int8 `json:"status"` Remark string `json:"remark"` } // clearTenantDefault 取消租户内所有默认配置 func clearTenantDefault(tenantID string) { var defaults []models.BackendAiProvider _, _ = models.Orm.QueryTable(new(models.BackendAiProvider)). Filter("tenant_id", tenantID). Filter("is_default", 1). All(&defaults) for _, d := range defaults { d.IsDefault = 0 d.UpdateTime = time.Now() _, _ = models.Orm.Update(&d, "is_default", "update_time") } } // parseModels 将模型JSON字符串解析为数组 func parseProviderModels(m string) []string { if m == "" { return []string{} } var list []string if err := json.Unmarshal([]byte(m), &list); err != nil { // 兼容旧格式:逗号分隔的单个模型 return []string{m} } return list } // List GET /backend/ai/provider/list func (c *BackendAiProviderController) List() { claims, err := c.aiClaims() if err != nil { c.aiJsonErr(401, 401, err.Error()) return } var list []models.BackendAiProvider _, err = models.Orm.QueryTable(new(models.BackendAiProvider)). Filter("tenant_id", fmt.Sprintf("%d", claims.TenantId)). Filter("user_id", uint64(claims.UserID)). Filter("delete_time__isnull", true). OrderBy("-id"). All(&list) if err != nil { c.aiJsonErr(500, 500, "查询失败: "+err.Error()) return } // 脱敏api_key + 解析模型列表 for i := range list { if list[i].ApiKey != "" && len(list[i].ApiKey) > 8 { list[i].ApiKey = list[i].ApiKey[:4] + "****" + list[i].ApiKey[len(list[i].ApiKey)-4:] } list[i].ModelsList = parseProviderModels(list[i].Models) } c.aiOk(map[string]interface{}{"list": list}) } // Create POST /backend/ai/provider func (c *BackendAiProviderController) Create() { claims, err := c.aiClaims() if err != nil { c.aiJsonErr(401, 401, err.Error()) return } body, err := io.ReadAll(c.Ctx.Request.Body) if err != nil { c.aiJsonErr(400, 400, "读取请求体失败") return } var p aiProviderPayload if err := json.Unmarshal(body, &p); err != nil { c.aiJsonErr(400, 400, "参数格式错误") return } if strings.TrimSpace(p.Name) == "" { c.aiJsonErr(400, 400, "配置名称不能为空") return } if p.ProviderType != "openai" && p.ProviderType != "anthropic" { c.aiJsonErr(400, 400, "接入类型必须是 openai 或 anthropic") return } if strings.TrimSpace(p.ApiBase) == "" || strings.TrimSpace(p.ApiKey) == "" { c.aiJsonErr(400, 400, "接口地址、API Key不能为空") return } if len(p.Models) == 0 { c.aiJsonErr(400, 400, "至少配置一个模型") return } // 模型数组序列化为JSON modelsJSON, _ := json.Marshal(p.Models) // 如果设为默认,先取消租户内其他默认 if p.IsDefault == 1 { clearTenantDefault(fmt.Sprintf("%d", claims.TenantId)) } provider := models.BackendAiProvider{ TenantID: fmt.Sprintf("%d", claims.TenantId), UserID: uint64(claims.UserID), ProviderType: p.ProviderType, Name: strings.TrimSpace(p.Name), ApiBase: strings.TrimSpace(p.ApiBase), ApiKey: strings.TrimSpace(p.ApiKey), Models: string(modelsJSON), IsDefault: p.IsDefault, Status: p.Status, Remark: strings.TrimSpace(p.Remark), CreateTime: time.Now(), UpdateTime: time.Now(), } id, err := models.Orm.Insert(&provider) if err != nil { c.aiJsonErr(500, 500, "创建失败: "+err.Error()) return } c.aiOk(map[string]interface{}{"id": id}) } // Update PUT /backend/ai/provider/:id func (c *BackendAiProviderController) Update() { claims, err := c.aiClaims() if err != nil { c.aiJsonErr(401, 401, err.Error()) return } idStr := c.Ctx.Input.Param(":id") id, err := strconv.ParseUint(idStr, 10, 64) if err != nil { c.aiJsonErr(400, 400, "ID格式错误") return } provider := models.BackendAiProvider{ID: id} if err := models.Orm.Read(&provider); err != nil { c.aiJsonErr(404, 404, "配置不存在") return } if provider.TenantID != fmt.Sprintf("%d", claims.TenantId) { c.aiJsonErr(403, 403, "无权操作") return } body, err := io.ReadAll(c.Ctx.Request.Body) if err != nil { c.aiJsonErr(400, 400, "读取请求体失败") return } var p aiProviderPayload if err := json.Unmarshal(body, &p); err != nil { c.aiJsonErr(400, 400, "参数格式错误") return } // 如果设为默认,先取消租户内其他默认 if p.IsDefault == 1 && provider.IsDefault != 1 { clearTenantDefault(fmt.Sprintf("%d", claims.TenantId)) } provider.Name = strings.TrimSpace(p.Name) provider.ProviderType = p.ProviderType provider.ApiBase = strings.TrimSpace(p.ApiBase) // 只有传入了非脱敏的key才更新 if p.ApiKey != "" && !strings.Contains(p.ApiKey, "****") { provider.ApiKey = strings.TrimSpace(p.ApiKey) } // 更新模型列表 if len(p.Models) > 0 { modelsJSON, _ := json.Marshal(p.Models) provider.Models = string(modelsJSON) } provider.IsDefault = p.IsDefault provider.Status = p.Status provider.Remark = strings.TrimSpace(p.Remark) provider.UpdateTime = time.Now() _, err = models.Orm.Update(&provider) if err != nil { c.aiJsonErr(500, 500, "更新失败: "+err.Error()) return } c.aiOk(nil) } // Delete DELETE /backend/ai/provider/:id func (c *BackendAiProviderController) Delete() { claims, err := c.aiClaims() if err != nil { c.aiJsonErr(401, 401, err.Error()) return } idStr := c.Ctx.Input.Param(":id") id, err := strconv.ParseUint(idStr, 10, 64) if err != nil { c.aiJsonErr(400, 400, "ID格式错误") return } provider := models.BackendAiProvider{ID: id} if err := models.Orm.Read(&provider); err != nil { c.aiJsonErr(404, 404, "配置不存在") return } if provider.TenantID != fmt.Sprintf("%d", claims.TenantId) { c.aiJsonErr(403, 403, "无权操作") return } now := time.Now() provider.DeleteTime = &now _, err = models.Orm.Update(&provider, "delete_time") if err != nil { c.aiJsonErr(500, 500, "删除失败: "+err.Error()) return } c.aiOk(nil) } type aiTestPayload struct { ProviderType string `json:"provider_type"` ApiBase string `json:"api_base"` ApiKey string `json:"api_key"` Models []string `json:"models"` } type modelTestResult struct { Model string `json:"model"` Success bool `json:"success"` Error string `json:"error"` } // Test POST /backend/ai/provider/test // 批量测试模型连通性 func (c *BackendAiProviderController) Test() { if _, err := c.aiClaims(); err != nil { c.aiJsonErr(401, 401, err.Error()) return } body, err := io.ReadAll(c.Ctx.Request.Body) if err != nil { c.aiJsonErr(400, 400, "读取请求体失败") return } var p aiTestPayload if err := json.Unmarshal(body, &p); err != nil { c.aiJsonErr(400, 400, "参数格式错误") return } if strings.TrimSpace(p.ApiBase) == "" || strings.TrimSpace(p.ApiKey) == "" { c.aiJsonErr(400, 400, "接口地址和API Key不能为空") return } if len(p.Models) == 0 { c.aiOk(map[string]interface{}{"results": []modelTestResult{}}) return } // 构造临时provider provider := models.BackendAiProvider{ ProviderType: p.ProviderType, ApiBase: strings.TrimSpace(p.ApiBase), ApiKey: strings.TrimSpace(p.ApiKey), } testMessages := []openaiMessage{{Role: "user", Content: "hi"}} // 并发测试 results := make([]modelTestResult, len(p.Models)) var wg sync.WaitGroup for i, model := range p.Models { wg.Add(1) go func(idx int, m string) { defer wg.Done() _, err := callAI(provider, m, "", testMessages) if err != nil { results[idx] = modelTestResult{Model: m, Success: false, Error: err.Error()} } else { results[idx] = modelTestResult{Model: m, Success: true, Error: ""} } }(i, model) } wg.Wait() c.aiOk(map[string]interface{}{"results": results}) }