109 lines
3.0 KiB
Go
109 lines
3.0 KiB
Go
package service
|
|
|
|
import (
|
|
"context"
|
|
|
|
"rag-local/kb/consts"
|
|
"rag-local/kb/dao"
|
|
"rag-local/kb/model/entity"
|
|
|
|
"github.com/cloudwego/eino/schema"
|
|
"github.com/gogf/gf/v2/errors/gerror"
|
|
"github.com/gogf/gf/v2/frame/g"
|
|
)
|
|
|
|
var ModelConfigService = &modelConfigService{}
|
|
|
|
type modelConfigService struct{}
|
|
|
|
func (s *modelConfigService) List(ctx context.Context, modelType string) ([]*entity.ModelConfig, error) {
|
|
return dao.ModelConfig.List(ctx, modelType)
|
|
}
|
|
|
|
func (s *modelConfigService) Save(ctx context.Context, m *entity.ModelConfig) (int64, error) {
|
|
// 留空的接口地址/维度按表中已有同名模型记录继承(用户显式填写的值优先,不覆盖)
|
|
if other, err := dao.ModelConfig.GetByName(ctx, m.ModelName, m.Id); err == nil && other != nil {
|
|
if m.EndpointUrl == "" && other.EndpointUrl != "" {
|
|
m.EndpointUrl = other.EndpointUrl
|
|
}
|
|
if m.ModelType == consts.ModelTypeEmbedding && m.Dimension <= 0 && other.Dimension > 0 {
|
|
m.Dimension = other.Dimension
|
|
}
|
|
}
|
|
if m.Dimension <= 0 {
|
|
m.Dimension = consts.DefaultEmbeddingDim
|
|
}
|
|
if m.Id > 0 {
|
|
if err := dao.ModelConfig.Update(ctx, m); err != nil {
|
|
return 0, err
|
|
}
|
|
return m.Id, nil
|
|
}
|
|
return dao.ModelConfig.Insert(ctx, m)
|
|
}
|
|
|
|
// SetDefault 设为同类型默认模型(同类型其余配置自动取消默认,保证唯一)
|
|
func (s *modelConfigService) SetDefault(ctx context.Context, id int64) error {
|
|
cfg, err := dao.ModelConfig.GetOne(ctx, id)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
if cfg == nil {
|
|
return gerror.New("模型配置不存在")
|
|
}
|
|
return dao.ModelConfig.SetDefault(ctx, id)
|
|
}
|
|
|
|
// Test 连通性测试:对话模型发一次 ping,向量模型嵌入一次,任一异常即失败
|
|
func (s *modelConfigService) Test(ctx context.Context, id int64) error {
|
|
cfg, err := dao.ModelConfig.GetOne(ctx, id)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
if cfg == nil {
|
|
return gerror.New("模型配置不存在")
|
|
}
|
|
switch cfg.ModelType {
|
|
case consts.ModelTypeChat:
|
|
model, err := BuildChatModel(ctx, id)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
msg, err := model.Generate(ctx, []*schema.Message{{Role: schema.User, Content: "ping"}})
|
|
if err != nil {
|
|
return gerror.Wrap(err, "对话接口调用失败")
|
|
}
|
|
if msg == nil {
|
|
return gerror.New("对话接口返回空")
|
|
}
|
|
case consts.ModelTypeEmbedding:
|
|
em, err := BuildEmbedder(ctx, id)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
vecs, err := em.EmbedStrings(ctx, []string{"ping"})
|
|
if err != nil {
|
|
return gerror.Wrap(err, "向量接口调用失败")
|
|
}
|
|
if len(vecs) == 0 || len(vecs[0]) == 0 {
|
|
return gerror.New("向量接口返回空数据")
|
|
}
|
|
default:
|
|
return gerror.New("未知模型类型")
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (s *modelConfigService) Delete(ctx context.Context, id int64) error {
|
|
// 被数据集绑定的 embedding 配置不允许删除
|
|
count, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameDataset).Ctx(ctx).
|
|
Where("embedding_cfg_id", id).Count()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
if count > 0 {
|
|
return gerror.New("该模型配置正被数据集使用,无法删除")
|
|
}
|
|
return dao.ModelConfig.Delete(ctx, id)
|
|
}
|