增加相关功能
This commit is contained in:
@@ -0,0 +1,360 @@
|
||||
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})
|
||||
}
|
||||
Reference in New Issue
Block a user