205 lines
5.4 KiB
Go
205 lines
5.4 KiB
Go
package service
|
|
|
|
import (
|
|
"context"
|
|
"fmt"
|
|
"model-gateway/consts/model"
|
|
"model-gateway/consts/public"
|
|
"model-gateway/dao"
|
|
"model-gateway/model/dto"
|
|
"model-gateway/model/entity"
|
|
"model-gateway/service/gateway"
|
|
|
|
"gitea.redpowerfuture.com/red-future/common/beans"
|
|
"gitea.redpowerfuture.com/red-future/common/db/gfdb"
|
|
"gitea.redpowerfuture.com/red-future/common/utils"
|
|
"github.com/gogf/gf/v2/database/gdb"
|
|
"github.com/gogf/gf/v2/frame/g"
|
|
"github.com/gogf/gf/v2/util/gconv"
|
|
)
|
|
|
|
var ModelManage = &modelManageService{}
|
|
|
|
type modelManageService struct{}
|
|
|
|
// Create 创建模型
|
|
func (s *modelManageService) Create(ctx context.Context, req *dto.CreateModelManageReq) (res *dto.CreateModelManageRes, err error) {
|
|
err = gfdb.DB(ctx, public.DbNameModelGateway).Transaction(ctx, func(ctx context.Context, tx gdb.TX) (err error) {
|
|
// 1)检查是否是超管
|
|
var isSuperAdmin bool
|
|
isSuperAdmin, err = gateway.IsSuperAdmin(ctx)
|
|
if err != nil {
|
|
return
|
|
}
|
|
req.SystemModel = &isSuperAdmin
|
|
// 1)如果设为会话模型,先把该用户旧会话模型取消
|
|
err = s.CancelChatModel(ctx, req.ModelType, req.ChatModel, isSuperAdmin)
|
|
if err != nil {
|
|
return
|
|
}
|
|
// 2)插入数据
|
|
id, err := dao.ModelManage.Insert(ctx, req)
|
|
if err != nil {
|
|
return
|
|
}
|
|
res = &dto.CreateModelManageRes{Id: id}
|
|
return
|
|
})
|
|
return
|
|
}
|
|
|
|
// Update 更新模型配置
|
|
func (s *modelManageService) Update(ctx context.Context, req *dto.UpdateModelManageReq) (err error) {
|
|
err = gfdb.DB(ctx, public.DbNameModelGateway).Transaction(ctx, func(ctx context.Context, tx gdb.TX) (err error) {
|
|
var get *entity.ModelManage
|
|
get, err = dao.ModelManage.GetNotTenantId(ctx, &dto.GetModelManageReq{
|
|
Id: req.Id,
|
|
})
|
|
if err != nil {
|
|
return
|
|
}
|
|
var user *beans.User
|
|
user, err = utils.GetUserInfo(ctx)
|
|
if err != nil {
|
|
return
|
|
}
|
|
// 1)如果不是创建者,且是系统模型,则需要拷贝
|
|
if get.Creator != user.UserName {
|
|
if get.SystemModel != nil && *get.SystemModel {
|
|
if g.IsEmpty(req.ApiKey) {
|
|
return fmt.Errorf("模型apiKey不能为空")
|
|
}
|
|
d := new(dto.CreateModelManageReq)
|
|
err = gconv.Struct(req, d)
|
|
if err != nil {
|
|
return
|
|
}
|
|
_, err = s.Create(ctx, d)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
return
|
|
}
|
|
return fmt.Errorf("无权限操作")
|
|
}
|
|
|
|
// 1)检查是否是超管
|
|
var isSuperAdmin bool
|
|
isSuperAdmin, err = gateway.IsSuperAdmin(ctx)
|
|
if err != nil {
|
|
return
|
|
}
|
|
// 1)如果设为会话模型,先把该用户旧会话模型取消
|
|
err = s.CancelChatModel(ctx, req.ModelType, req.ChatModel, isSuperAdmin)
|
|
if err != nil {
|
|
return
|
|
}
|
|
// 2)更新数据
|
|
_, err = dao.ModelManage.Update(ctx, req)
|
|
return
|
|
})
|
|
return
|
|
}
|
|
|
|
func (s *modelManageService) CancelChatModel(ctx context.Context, modelType model.ModelType, chatModel *bool, isSuperAdmin bool) (err error) {
|
|
if !g.IsEmpty(chatModel) && *chatModel {
|
|
if *modelType == *model.ModelTypeInference.Code {
|
|
if isSuperAdmin {
|
|
return fmt.Errorf("超级管理员不能设置会话模型")
|
|
}
|
|
// 2)获取该用户信息
|
|
var user *beans.User
|
|
user, err = utils.GetUserInfo(ctx)
|
|
if err != nil {
|
|
return
|
|
}
|
|
// 3)取消该用户之前的会话模型
|
|
var get *entity.ModelManage
|
|
get, err = dao.ModelManage.Get(ctx, &dto.GetModelManage{
|
|
Creator: user.UserName,
|
|
ChatModel: chatModel,
|
|
})
|
|
if err != nil {
|
|
return
|
|
}
|
|
_, err = dao.ModelManage.Update(ctx, &dto.UpdateModelManageReq{
|
|
Id: get.Id,
|
|
ChatModel: gconv.PtrBool(false),
|
|
})
|
|
if err != nil {
|
|
return
|
|
}
|
|
} else {
|
|
return fmt.Errorf("只有推理模型可以设置成会话模型")
|
|
}
|
|
}
|
|
return
|
|
}
|
|
|
|
// Delete 删除模型
|
|
func (s *modelManageService) Delete(ctx context.Context, req *dto.DeleteModelManageReq) error {
|
|
_, err := dao.ModelManage.Delete(ctx, req)
|
|
return err
|
|
}
|
|
|
|
func (s *modelManageService) Get(ctx context.Context, req *dto.GetModelManageReq) (res *dto.GetModelManageRes, err error) {
|
|
get, err := dao.ModelManage.GetNotTenantId(ctx, req)
|
|
if err != nil {
|
|
return
|
|
}
|
|
err = gconv.Struct(get, &res)
|
|
return
|
|
}
|
|
|
|
// List 获取模型列表
|
|
func (s *modelManageService) List(ctx context.Context, req *dto.ListModelManageReq) (res *dto.ListModelManageRes, err error) {
|
|
var user *beans.User
|
|
user, err = utils.GetUserInfo(ctx)
|
|
if err != nil {
|
|
return
|
|
}
|
|
req.Creator = user.UserName
|
|
list, total, err := dao.ModelManage.ListNotTenantId(ctx, req)
|
|
if err != nil {
|
|
return
|
|
}
|
|
res = &dto.ListModelManageRes{
|
|
Total: total,
|
|
}
|
|
err = gconv.Struct(list, &res.List)
|
|
return
|
|
}
|
|
|
|
func (s *modelManageService) CheckChatModel(ctx context.Context, req *dto.CheckChatModelReq) (res *dto.CheckChatModelRes, err error) {
|
|
user, err := utils.GetUserInfo(ctx)
|
|
if err != nil {
|
|
return
|
|
}
|
|
get, err := dao.ModelManage.Get(ctx, &dto.GetModelManage{
|
|
Creator: user.UserName,
|
|
ChatModel: gconv.PtrBool(true),
|
|
})
|
|
if err != nil {
|
|
return
|
|
}
|
|
res = &dto.CheckChatModelRes{
|
|
IsChatModel: !g.IsEmpty(get),
|
|
}
|
|
return
|
|
}
|
|
|
|
// GetModelType 获取模型类型
|
|
func (s *modelManageService) GetModelType(ctx context.Context, req *dto.ModelTypeReq) (res *dto.ModelTypeRes, err error) {
|
|
res = &dto.ModelTypeRes{
|
|
List: model.GetTypeTreeList(),
|
|
}
|
|
return res, nil
|
|
}
|
|
|
|
// GetModelSupplier 获取运营商列表
|
|
func (s *modelManageService) GetModelSupplier(ctx context.Context, req *dto.ModelSupplierReq) (res *dto.ModelSupplierRes, err error) {
|
|
return &dto.ModelSupplierRes{
|
|
List: model.GetSupplierOptionList(),
|
|
}, nil
|
|
}
|