From b11ed5b433f9c75b1f2146ff98eff92e55b6a245 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=BC=A0=E6=96=8C?= <259278618@qq.com> Date: Wed, 5 Aug 2026 13:39:37 +0800 Subject: [PATCH] 1 --- common/auth.go | 16 +++ common/auth_middleware.go | 8 +- config.yml | 6 +- go.mod | 2 + go.sum | 4 + kb/consts/consts.go | 8 +- kb/consts/content_type.go | 7 - kb/consts/table_name.go | 1 - kb/controller/dataset_controller.go | 3 + kb/controller/model_config_controller.go | 4 + kb/controller/system_config_controller.go | 33 ----- kb/dao/chunk_dao.go | 3 + kb/dao/conversation_dao.go | 3 + kb/dao/dataset_dao.go | 32 +++++ kb/dao/document_dao.go | 3 + kb/dao/kg_entity_dao.go | 3 + kb/dao/kg_relation_dao.go | 3 + kb/dao/message_dao.go | 3 + kb/dao/model_config_dao.go | 55 +++++++ kb/dao/parse_task_dao.go | 3 + kb/dao/system_config_dao.go | 57 -------- kb/model/dto/dataset_dto.go | 3 + kb/model/dto/model_config_dto.go | 7 + kb/model/dto/system_config_dto.go | 33 ----- kb/model/entity/dataset.go | 3 + kb/model/entity/model_config.go | 1 + kb/model/entity/system_config.go | 10 -- kb/service/chat_service.go | 55 ++++--- kb/service/chunk_service.go | 83 +++++++++-- kb/service/dataset_service.go | 55 +++++++ kb/service/kg_entity_service.go | 5 +- kb/service/model_config_service.go | 12 ++ kb/service/parse_task_service.go | 61 ++++++-- kb/service/system_config_service.go | 83 +---------- ui-src/src/api/auth.js | 16 --- ui-src/src/api/model_config.js | 4 + ui-src/src/views/Chat.vue | 6 +- ui-src/src/views/DatasetDetail.vue | 2 +- ui-src/src/views/DatasetList.vue | 74 +++++++++- ui-src/src/views/Settings.vue | 167 +++++----------------- 40 files changed, 503 insertions(+), 434 deletions(-) delete mode 100644 kb/dao/system_config_dao.go delete mode 100644 kb/model/entity/system_config.go diff --git a/common/auth.go b/common/auth.go index 02f48f7..54e2820 100644 --- a/common/auth.go +++ b/common/auth.go @@ -32,6 +32,22 @@ func SignToken(role, tokenFp string, expireSeconds int64) (string, error) { return jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString([]byte(jwtSecret)) } +// ---------- 访问令牌(内存持有,每次启动重新生成,不落库) ---------- + +var accessToken string + +func SetAccessToken(token string) { + accessToken = token +} + +func CheckAccessToken(token string) bool { + return accessToken != "" && token == accessToken +} + +func AccessTokenFingerprint() string { + return TokenFingerprint(accessToken) +} + func ParseToken(tokenStr string) (*JwtClaims, error) { token, err := jwt.ParseWithClaims(tokenStr, &JwtClaims{}, func(token *jwt.Token) (interface{}, error) { return []byte(jwtSecret), nil diff --git a/common/auth_middleware.go b/common/auth_middleware.go index 573cc29..2ceb9ed 100644 --- a/common/auth_middleware.go +++ b/common/auth_middleware.go @@ -1,7 +1,6 @@ package common import ( - "context" "net/http" "strings" @@ -12,9 +11,6 @@ var publicPaths = []string{ "/system-config/login", } -// CheckTokenFingerprint 由 kb/service 注入,避免 common → service 循环依赖 -var CheckTokenFingerprint func(ctx context.Context, fp string) bool - func Auth(r *ghttp.Request) { path := r.URL.Path @@ -58,8 +54,8 @@ func Auth(r *ghttp.Request) { return } - // 指纹校验:访问令牌被重新生成后,旧会话立即失效 - if CheckTokenFingerprint != nil && !CheckTokenFingerprint(r.Context(), claims.TokenFp) { + // 指纹校验:访问令牌变更(重启)后,旧会话立即失效 + if claims.TokenFp != AccessTokenFingerprint() { r.Response.WriteJson(ghttp.DefaultHandlerResponse{ Code: http.StatusUnauthorized, Message: "访问令牌已变更,请重新登录", diff --git a/config.yml b/config.yml index e27927f..04ef6ae 100644 --- a/config.yml +++ b/config.yml @@ -2,15 +2,15 @@ database: default: name: data/business.db type: sqlite - debug: false + debug: true # 控制台打印所有 SQL 查询 system: name: data/system.db type: sqlite - debug: false + debug: true # 控制台打印所有 SQL 查询 chat: name: data/chat.db type: sqlite - debug: false + debug: true # 控制台打印所有 SQL 查询 cache: ttl: 60 # DAO查询缓存时间(秒),0为禁用缓存 server: diff --git a/go.mod b/go.mod index 18a6051..ef0afcd 100644 --- a/go.mod +++ b/go.mod @@ -5,6 +5,8 @@ go 1.26.1 require ( github.com/PuerkitoBio/goquery v1.12.0 github.com/cloudwego/eino v0.9.13 + github.com/cloudwego/eino-ext/components/document/transformer/splitter/recursive v0.0.0-20260803030130-90a15623ddb6 + github.com/cloudwego/eino-ext/components/document/transformer/splitter/semantic v0.0.0-20260803030130-90a15623ddb6 github.com/go-ego/gse v1.0.2 github.com/gogf/gf/contrib/drivers/sqlite/v2 v2.10.2 github.com/gogf/gf/v2 v2.10.2 diff --git a/go.sum b/go.sum index 63af4be..1a5e65d 100644 --- a/go.sum +++ b/go.sum @@ -28,6 +28,10 @@ github.com/cloudwego/base64x v0.1.6 h1:t11wG9AECkCDk5fMSoxmufanudBtJ+/HemLstXDLI github.com/cloudwego/base64x v0.1.6/go.mod h1:OFcloc187FXDaYHvrNIjxSe8ncn0OOM8gEHfghB2IPU= github.com/cloudwego/eino v0.9.13 h1:iD/ETS+lxnNp1VeNPqWVGPWdND6Dbf4LyINbLUlDRcM= github.com/cloudwego/eino v0.9.13/go.mod h1:OBD1mrkfkt/pJa4rkg1P0VnaMeOVl7l8IAdEqY//3IQ= +github.com/cloudwego/eino-ext/components/document/transformer/splitter/recursive v0.0.0-20260803030130-90a15623ddb6 h1:jQxOBnesRjdFQrIcZ5l4BZ3RvcY2La3/V+GUzAtnJv8= +github.com/cloudwego/eino-ext/components/document/transformer/splitter/recursive v0.0.0-20260803030130-90a15623ddb6/go.mod h1:9R0RQrQSpg1JaNnRtw7+RfRAAv0HgdE348YnrlZ6coo= +github.com/cloudwego/eino-ext/components/document/transformer/splitter/semantic v0.0.0-20260803030130-90a15623ddb6 h1:tbm/a5vzwVv2YevFcJgXiaHj9MNFnyYpMd7zC75ACkM= +github.com/cloudwego/eino-ext/components/document/transformer/splitter/semantic v0.0.0-20260803030130-90a15623ddb6/go.mod h1:Ov33JMUewdOoUgJbYNJt3qL7KQDVHYpoVBCjJXsz8sw= github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= diff --git a/kb/consts/consts.go b/kb/consts/consts.go index 0d56e1e..9e74358 100644 --- a/kb/consts/consts.go +++ b/kb/consts/consts.go @@ -8,8 +8,12 @@ const ( HybridTopK = 10 // 混合检索融合后返回数 RrfK = 60 // RRF 融合常数 - MaxChunkSize = 800 // 分块最大字数 - ChunkOverlap = 100 // 分块重叠字数 + DefaultChunkSize = 800 // 分块最大字数(数据集默认值) + DefaultChunkOverlap = 150 // 分块重叠字数(数据集默认值) + + ChunkStrategyTitle = "title" // 标题感知分块(自研,保留标题与段落) + ChunkStrategyRecursive = "recursive" // 递归字符分块(Eino recursive) + ChunkStrategySemantic = "semantic" // 语义分块(Eino semantic,需绑定向量模型) ParsePollIntervalSeconds = 3 // 解析任务轮询间隔 diff --git a/kb/consts/content_type.go b/kb/consts/content_type.go index 98ba1e2..418b12d 100644 --- a/kb/consts/content_type.go +++ b/kb/consts/content_type.go @@ -5,10 +5,3 @@ const ( ModelTypeChat = "chat" ModelTypeEmbedding = "embedding" ) - -// system_config 预置键 -const ( - CfgKeyAccessToken = "access_token" - CfgKeyDefaultChatModel = "default_chat_model" - CfgKeyDefaultDataset = "default_dataset" -) diff --git a/kb/consts/table_name.go b/kb/consts/table_name.go index b8485ea..4465726 100644 --- a/kb/consts/table_name.go +++ b/kb/consts/table_name.go @@ -1,7 +1,6 @@ package consts const ( - TableNameSystemConfig = "system_config" TableNameModelConfig = "model_config" TableNameDataset = "kb_dataset" TableNameDocument = "kb_document" diff --git a/kb/controller/dataset_controller.go b/kb/controller/dataset_controller.go index d6da4fe..9820ebb 100644 --- a/kb/controller/dataset_controller.go +++ b/kb/controller/dataset_controller.go @@ -26,6 +26,9 @@ func (c *dataset) Save(ctx context.Context, req *dto.SaveDatasetReq) (*dto.SaveD Name: req.Name, Description: req.Description, EmbeddingCfgId: req.EmbeddingCfgId, + ChunkSize: req.ChunkSize, + ChunkOverlap: req.ChunkOverlap, + ChunkStrategy: req.ChunkStrategy, Status: 1, }) if err != nil { diff --git a/kb/controller/model_config_controller.go b/kb/controller/model_config_controller.go index 5efd8d8..830bc13 100644 --- a/kb/controller/model_config_controller.go +++ b/kb/controller/model_config_controller.go @@ -37,6 +37,10 @@ func (c *modelConfig) Save(ctx context.Context, req *dto.SaveModelConfigReq) (re return &dto.SaveModelConfigRes{Id: id}, nil } +func (c *modelConfig) SetDefault(ctx context.Context, req *dto.SetDefaultModelConfigReq) (res *dto.SetDefaultModelConfigRes, err error) { + return nil, service.ModelConfigService.SetDefault(ctx, req.Id) +} + func (c *modelConfig) Delete(ctx context.Context, req *dto.DeleteModelConfigReq) (res *dto.DeleteModelConfigRes, err error) { return nil, service.ModelConfigService.Delete(ctx, req.Id) } diff --git a/kb/controller/system_config_controller.go b/kb/controller/system_config_controller.go index d3aaac2..b6550e6 100644 --- a/kb/controller/system_config_controller.go +++ b/kb/controller/system_config_controller.go @@ -2,7 +2,6 @@ package controller import ( "context" - "fmt" "rag-local/kb/model/dto" "rag-local/kb/service" @@ -19,35 +18,3 @@ func (c *systemConfig) Login(ctx context.Context, req *dto.LoginReq) (res *dto.L } return &dto.LoginRes{Token: token}, nil } - -func (c *systemConfig) Get(ctx context.Context, req *dto.GetSystemConfigReq) (res *dto.GetSystemConfigRes, err error) { - chatModel, dataset, err := service.SystemConfigService.GetSettings(ctx) - if err != nil { - return nil, err - } - return &dto.GetSystemConfigRes{DefaultChatModel: chatModel, DefaultDataset: dataset}, nil -} - -func (c *systemConfig) Update(ctx context.Context, req *dto.UpdateSystemConfigReq) (res *dto.UpdateSystemConfigRes, err error) { - return nil, service.SystemConfigService.UpdateSettings(ctx, req.DefaultChatModel, req.DefaultDataset) -} - -func (c *systemConfig) GetToken(ctx context.Context, req *dto.GetTokenReq) (*dto.GetTokenRes, error) { - token, err := service.SystemConfigService.GetToken(ctx) - if err != nil { - return nil, err - } - return &dto.GetTokenRes{Token: token}, nil -} - -func (c *systemConfig) RegenerateToken(ctx context.Context, req *dto.RegenerateTokenReq) (res *dto.RegenerateTokenRes, err error) { - token, err := service.SystemConfigService.RegenerateToken(ctx) - if err != nil { - return nil, err - } - fmt.Printf("\n============================================\n") - fmt.Printf("访问令牌(登录用)已更新: %s\n", token) - fmt.Printf("旧令牌已失效,请在登录页重新输入\n") - fmt.Printf("============================================\n\n") - return &dto.RegenerateTokenRes{Token: token}, nil -} diff --git a/kb/dao/chunk_dao.go b/kb/dao/chunk_dao.go index 4bf20f1..a20ecd5 100644 --- a/kb/dao/chunk_dao.go +++ b/kb/dao/chunk_dao.go @@ -83,6 +83,9 @@ func (d *chunkDao) ListByDocument(ctx context.Context, documentId int64, page, p var list []*entity.Chunk err = g.DB(consts.DbGroupDefault).Model(consts.TableNameChunk).Ctx(ctx). Where("document_id", documentId).Page(page, pageSize).OrderAsc("seq").Scan(&list) + if list == nil { + list = make([]*entity.Chunk, 0) + } return list, total, err } diff --git a/kb/dao/conversation_dao.go b/kb/dao/conversation_dao.go index ca11ed7..3ef22e7 100644 --- a/kb/dao/conversation_dao.go +++ b/kb/dao/conversation_dao.go @@ -45,6 +45,9 @@ func (d *conversationDao) GetOne(ctx context.Context, id int64) (*entity.Convers func (d *conversationDao) List(ctx context.Context) ([]*entity.Conversation, error) { var list []*entity.Conversation err := g.DB(consts.DbGroupChat).Model(consts.TableNameConversation).Ctx(ctx).OrderDesc("id").Scan(&list) + if list == nil { + list = make([]*entity.Conversation, 0) + } return list, err } diff --git a/kb/dao/dataset_dao.go b/kb/dao/dataset_dao.go index 07aa76f..1358bc6 100644 --- a/kb/dao/dataset_dao.go +++ b/kb/dao/dataset_dao.go @@ -23,6 +23,9 @@ func init() { name TEXT NOT NULL DEFAULT '', description TEXT NOT NULL DEFAULT '', embedding_cfg_id INTEGER NOT NULL DEFAULT 0, + chunk_size INTEGER NOT NULL DEFAULT 800, + chunk_overlap INTEGER NOT NULL DEFAULT 150, + chunk_strategy TEXT NOT NULL DEFAULT 'title', status INTEGER NOT NULL DEFAULT 1, created_at DATETIME DEFAULT (datetime('now','localtime')), updated_at DATETIME DEFAULT (datetime('now','localtime')) @@ -30,6 +33,26 @@ func init() { if err != nil { g.Log().Warningf(ctx, "create kb_dataset table failed: %v", err) } + // 旧表迁移:补充分块配置列(先查列是否已存在,避免 ALTER 报错刷日志) + for _, col := range []struct { + name string + ddl string + }{ + {"chunk_size", "chunk_size INTEGER NOT NULL DEFAULT 800"}, + {"chunk_overlap", "chunk_overlap INTEGER NOT NULL DEFAULT 150"}, + {"chunk_strategy", "chunk_strategy TEXT NOT NULL DEFAULT 'title'"}, + } { + var n int + err := g.DB(consts.DbGroupDefault).Ctx(ctx).GetScan(ctx, &n, + "SELECT COUNT(*) FROM pragma_table_info('"+consts.TableNameDataset+"') WHERE name=?", col.name) + if err != nil || n > 0 { + continue + } + if _, err := g.DB(consts.DbGroupDefault).Exec(ctx, + "ALTER TABLE "+consts.TableNameDataset+" ADD COLUMN "+col.ddl); err != nil { + g.Log().Warningf(ctx, "migrate kb_dataset add column %s failed: %v", col.name, err) + } + } } func (d *datasetDao) GetOne(ctx context.Context, id int64) (*entity.Dataset, error) { @@ -47,6 +70,9 @@ func (d *datasetDao) GetOne(ctx context.Context, id int64) (*entity.Dataset, err func (d *datasetDao) List(ctx context.Context) ([]*entity.Dataset, error) { var list []*entity.Dataset err := g.DB(consts.DbGroupDefault).Model(consts.TableNameDataset).Ctx(ctx).OrderAsc("id").Scan(&list) + if list == nil { + list = make([]*entity.Dataset, 0) + } return list, err } @@ -56,6 +82,9 @@ func (d *datasetDao) Insert(ctx context.Context, data *entity.Dataset) (int64, e "name": data.Name, "description": data.Description, "embedding_cfg_id": data.EmbeddingCfgId, + "chunk_size": data.ChunkSize, + "chunk_overlap": data.ChunkOverlap, + "chunk_strategy": data.ChunkStrategy, "status": data.Status, "created_at": now, "updated_at": now, @@ -71,6 +100,9 @@ func (d *datasetDao) Update(ctx context.Context, data *entity.Dataset) error { "name": data.Name, "description": data.Description, "embedding_cfg_id": data.EmbeddingCfgId, + "chunk_size": data.ChunkSize, + "chunk_overlap": data.ChunkOverlap, + "chunk_strategy": data.ChunkStrategy, "status": data.Status, "updated_at": gtime.Now().Format("Y-m-d H:i:s"), }).Where("id", data.Id).Update() diff --git a/kb/dao/document_dao.go b/kb/dao/document_dao.go index f879529..2afbd83 100644 --- a/kb/dao/document_dao.go +++ b/kb/dao/document_dao.go @@ -65,6 +65,9 @@ func (d *documentDao) List(ctx context.Context, datasetId int64, page, pageSize var list []*entity.Document err = g.DB(consts.DbGroupDefault).Model(consts.TableNameDocument).Ctx(ctx). Where("dataset_id", datasetId).Page(page, pageSize).OrderDesc("id").Scan(&list) + if list == nil { + list = make([]*entity.Document, 0) + } return list, total, err } diff --git a/kb/dao/kg_entity_dao.go b/kb/dao/kg_entity_dao.go index ff2b156..3035d47 100644 --- a/kb/dao/kg_entity_dao.go +++ b/kb/dao/kg_entity_dao.go @@ -66,6 +66,9 @@ func (d *kgEntityDao) List(ctx context.Context, datasetId int64, page, pageSize } var list []*entity.KgEntity err = m.Page(page, pageSize).OrderDesc("id").Scan(&list) + if list == nil { + list = make([]*entity.KgEntity, 0) + } return list, total, err } diff --git a/kb/dao/kg_relation_dao.go b/kb/dao/kg_relation_dao.go index d3e68be..3265bb4 100644 --- a/kb/dao/kg_relation_dao.go +++ b/kb/dao/kg_relation_dao.go @@ -65,6 +65,9 @@ func (d *kgRelationDao) List(ctx context.Context, datasetId int64, page, pageSiz } var list []*entity.KgRelation err = m.Page(page, pageSize).OrderDesc("id").Scan(&list) + if list == nil { + list = make([]*entity.KgRelation, 0) + } return list, total, err } diff --git a/kb/dao/message_dao.go b/kb/dao/message_dao.go index 041c90a..aef3ec8 100644 --- a/kb/dao/message_dao.go +++ b/kb/dao/message_dao.go @@ -36,6 +36,9 @@ func (d *messageDao) List(ctx context.Context, conversationId int64) ([]*entity. var list []*entity.Message err := g.DB(consts.DbGroupChat).Model(consts.TableNameMessage).Ctx(ctx). Where("conversation_id", conversationId).OrderAsc("id").Scan(&list) + if list == nil { + list = make([]*entity.Message, 0) + } return list, err } diff --git a/kb/dao/model_config_dao.go b/kb/dao/model_config_dao.go index 1e52847..1850ed6 100644 --- a/kb/dao/model_config_dao.go +++ b/kb/dao/model_config_dao.go @@ -30,12 +30,35 @@ func init() { api_key TEXT NOT NULL DEFAULT '', dimension INTEGER NOT NULL DEFAULT 1024, extra TEXT NOT NULL DEFAULT '', + is_default INTEGER NOT NULL DEFAULT 0, created_at DATETIME DEFAULT (datetime('now','localtime')), updated_at DATETIME DEFAULT (datetime('now','localtime')) )`) if err != nil { g.Log().Warningf(ctx, "create model_config table failed: %v", err) } + // 迁移:旧库补 is_default 列,并把 system_config 中的旧默认对话模型键迁移为 is_default + var n int + if err := g.DB(consts.DbGroupSystem).Ctx(ctx).GetScan(ctx, &n, + "SELECT COUNT(*) FROM pragma_table_info('"+consts.TableNameModelConfig+"') WHERE name=?", "is_default"); err == nil && n == 0 { + if _, err := g.DB(consts.DbGroupSystem).Exec(ctx, "ALTER TABLE "+consts.TableNameModelConfig+" ADD COLUMN is_default INTEGER NOT NULL DEFAULT 0"); err != nil { + g.Log().Warningf(ctx, "alter model_config add is_default failed: %v", err) + } else { + var old int64 + if err := g.DB(consts.DbGroupSystem).Ctx(ctx).GetScan(ctx, &old, + "SELECT cfg_value FROM system_config WHERE cfg_key='default_chat_model'"); err == nil && old > 0 { + if _, err := g.DB(consts.DbGroupSystem).Exec(ctx, + "UPDATE "+consts.TableNameModelConfig+" SET is_default=CASE WHEN id=? THEN 1 ELSE 0 END WHERE model_type=?", + old, consts.ModelTypeChat); err != nil { + g.Log().Warningf(ctx, "migrate default_chat_model to is_default failed: %v", err) + } + } + } + } + // system_config 表已废弃(令牌改内存、默认模型迁入 is_default),清理残留 + if _, err := g.DB(consts.DbGroupSystem).Exec(ctx, "DROP TABLE IF EXISTS system_config"); err != nil { + g.Log().Warningf(ctx, "drop system_config table failed: %v", err) + } } func (d *modelConfigDao) GetOne(ctx context.Context, id int64) (*entity.ModelConfig, error) { @@ -80,11 +103,21 @@ func (d *modelConfigDao) List(ctx context.Context, modelType string) ([]*entity. m = m.Where("model_type", modelType) } err := m.Scan(&list) + if list == nil { + list = make([]*entity.ModelConfig, 0) + } return list, err } func (d *modelConfigDao) Insert(ctx context.Context, data *entity.ModelConfig) (int64, error) { now := gtime.Now().Format("Y-m-d H:i:s") + if data.IsDefault == 1 { + // 新配置设为默认时,同类型旧默认先清零,保证唯一 + if _, err := g.DB(consts.DbGroupSystem).Model(consts.TableNameModelConfig).Ctx(ctx). + Where("model_type", data.ModelType).Update(g.Map{"is_default": 0}); err != nil { + return 0, err + } + } r, err := g.DB(consts.DbGroupSystem).Model(consts.TableNameModelConfig).Ctx(ctx).Data(g.Map{ "name": data.Name, "model_type": data.ModelType, @@ -93,6 +126,7 @@ func (d *modelConfigDao) Insert(ctx context.Context, data *entity.ModelConfig) ( "api_key": data.ApiKey, "dimension": data.Dimension, "extra": data.Extra, + "is_default": data.IsDefault, "created_at": now, "updated_at": now, }).Insert() @@ -102,6 +136,27 @@ func (d *modelConfigDao) Insert(ctx context.Context, data *entity.ModelConfig) ( return r.LastInsertId() } +// GetDefault 返回指定类型的默认模型配置 id;无默认返回 0 +func (d *modelConfigDao) GetDefault(ctx context.Context, modelType string) (int64, error) { + v, err := g.DB(consts.DbGroupSystem).Model(consts.TableNameModelConfig).Ctx(ctx). + Where("model_type", modelType).Where("is_default", 1).Value("id") + if err != nil { + return 0, err + } + if v.IsEmpty() { + return 0, nil + } + return v.Int64(), nil +} + +// SetDefault 设置 id 为同类型默认,其余同类型取消默认(单 SQL 保证唯一) +func (d *modelConfigDao) SetDefault(ctx context.Context, id int64) error { + _, err := g.DB(consts.DbGroupSystem).Exec(ctx, + "UPDATE "+consts.TableNameModelConfig+" SET is_default=CASE WHEN id=? THEN 1 ELSE 0 END WHERE model_type=(SELECT model_type FROM "+consts.TableNameModelConfig+" WHERE id=?)", + id, id) + return err +} + func (d *modelConfigDao) Update(ctx context.Context, data *entity.ModelConfig) error { _, err := g.DB(consts.DbGroupSystem).Model(consts.TableNameModelConfig).Ctx(ctx).Data(g.Map{ "name": data.Name, diff --git a/kb/dao/parse_task_dao.go b/kb/dao/parse_task_dao.go index c54018f..3f2dbee 100644 --- a/kb/dao/parse_task_dao.go +++ b/kb/dao/parse_task_dao.go @@ -62,6 +62,9 @@ func (d *parseTaskDao) List(ctx context.Context, page, pageSize int) ([]*entity. var list []*entity.ParseTask err = g.DB(consts.DbGroupDefault).Model(consts.TableNameParseTask).Ctx(ctx). Page(page, pageSize).OrderDesc("id").Scan(&list) + if list == nil { + list = make([]*entity.ParseTask, 0) + } return list, total, err } diff --git a/kb/dao/system_config_dao.go b/kb/dao/system_config_dao.go deleted file mode 100644 index b24ebf6..0000000 --- a/kb/dao/system_config_dao.go +++ /dev/null @@ -1,57 +0,0 @@ -package dao - -import ( - "context" - - "rag-local/common" - "rag-local/kb/consts" - - "github.com/gogf/gf/v2/database/gdb" - "github.com/gogf/gf/v2/frame/g" -) - -var SystemConfig = &systemConfigDao{} - -type systemConfigDao struct{} - -func init() { - ctx := context.Background() - _, err := g.DB(consts.DbGroupSystem).Exec(ctx, `CREATE TABLE IF NOT EXISTS `+consts.TableNameSystemConfig+` ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - cfg_key TEXT NOT NULL DEFAULT '', - cfg_value TEXT NOT NULL DEFAULT '', - updated_at DATETIME DEFAULT (datetime('now','localtime')) - )`) - if err != nil { - g.Log().Warningf(ctx, "create system_config table failed: %v", err) - } - if _, err := g.DB(consts.DbGroupSystem).Exec(ctx, "CREATE UNIQUE INDEX IF NOT EXISTS idx_system_config_key ON "+consts.TableNameSystemConfig+"(cfg_key)"); err != nil { - g.Log().Warningf(ctx, "create index idx_system_config_key failed: %v", err) - } -} - -func (d *systemConfigDao) Get(ctx context.Context, key string) (string, error) { - r, err := g.DB(consts.DbGroupSystem).Model(consts.TableNameSystemConfig).Ctx(ctx). - Cache(gdb.CacheOption{Duration: common.CacheTTL(), Name: "system_config_Get_" + key}). - Fields("cfg_value").Where("cfg_key", key).One() - if err != nil { - return "", err - } - if r == nil { - return "", nil - } - return r["cfg_value"].String(), nil -} - -func (d *systemConfigDao) Set(ctx context.Context, key, value string) error { - _, err := g.DB(consts.DbGroupSystem).Exec(ctx, - "INSERT INTO "+consts.TableNameSystemConfig+" (cfg_key, cfg_value, updated_at) VALUES (?, ?, datetime('now','localtime')) "+ - "ON CONFLICT(cfg_key) DO UPDATE SET cfg_value=excluded.cfg_value, updated_at=datetime('now','localtime')", - key, value) - if err != nil { - return err - } - // 查询缓存挂在 DB 内部缓存实例上且键带 SelectCache: 前缀,必须用 GetCache() 按完整键移除 - _, _ = g.DB(consts.DbGroupSystem).GetCache().Remove(ctx, "SelectCache:system_config_Get_"+key) - return nil -} diff --git a/kb/model/dto/dataset_dto.go b/kb/model/dto/dataset_dto.go index 536d0fa..077b7bc 100644 --- a/kb/model/dto/dataset_dto.go +++ b/kb/model/dto/dataset_dto.go @@ -20,6 +20,9 @@ type SaveDatasetReq struct { Name string `v:"required" json:"name"` Description string `json:"description"` EmbeddingCfgId int64 `json:"embedding_cfg_id"` + ChunkSize int `json:"chunk_size"` + ChunkOverlap int `json:"chunk_overlap"` + ChunkStrategy string `json:"chunk_strategy"` } type SaveDatasetRes struct { diff --git a/kb/model/dto/model_config_dto.go b/kb/model/dto/model_config_dto.go index 49af794..e2d04be 100644 --- a/kb/model/dto/model_config_dto.go +++ b/kb/model/dto/model_config_dto.go @@ -38,6 +38,13 @@ type DeleteModelConfigReq struct { type DeleteModelConfigRes struct{} +type SetDefaultModelConfigReq struct { + g.Meta `path:"/set-default" method:"post" tags:"模型配置" summary:"设为默认模型"` + Id int64 `v:"required" json:"id"` +} + +type SetDefaultModelConfigRes struct{} + type TestModelConfigReq struct { g.Meta `path:"/test" method:"post" tags:"模型配置" summary:"连通性测试"` Id int64 `v:"required" json:"id"` diff --git a/kb/model/dto/system_config_dto.go b/kb/model/dto/system_config_dto.go index 34199e7..44b74d7 100644 --- a/kb/model/dto/system_config_dto.go +++ b/kb/model/dto/system_config_dto.go @@ -10,36 +10,3 @@ type LoginReq struct { type LoginRes struct { Token string `json:"token"` } - -type GetSystemConfigReq struct { - g.Meta `path:"/" method:"get" tags:"系统配置" summary:"获取系统设置"` -} - -type GetSystemConfigRes struct { - DefaultChatModel int64 `json:"default_chat_model"` - DefaultDataset int64 `json:"default_dataset"` -} - -type UpdateSystemConfigReq struct { - g.Meta `path:"/" method:"put" tags:"系统配置" summary:"更新系统设置"` - DefaultChatModel int64 `json:"default_chat_model"` - DefaultDataset int64 `json:"default_dataset"` -} - -type UpdateSystemConfigRes struct{} - -type RegenerateTokenReq struct { - g.Meta `path:"/regenerate-token" method:"post" tags:"系统配置" summary:"重新生成访问令牌"` -} - -type RegenerateTokenRes struct { - Token string `json:"token"` -} - -type GetTokenReq struct { - g.Meta `path:"/token" method:"get" tags:"系统配置" summary:"获取当前访问令牌"` -} - -type GetTokenRes struct { - Token string `json:"token"` -} diff --git a/kb/model/entity/dataset.go b/kb/model/entity/dataset.go index 2505ac6..3568993 100644 --- a/kb/model/entity/dataset.go +++ b/kb/model/entity/dataset.go @@ -7,6 +7,9 @@ type Dataset struct { Name string `orm:"name" json:"name"` Description string `orm:"description" json:"description"` EmbeddingCfgId int64 `orm:"embedding_cfg_id" json:"embedding_cfg_id"` + ChunkSize int `orm:"chunk_size" json:"chunk_size"` + ChunkOverlap int `orm:"chunk_overlap" json:"chunk_overlap"` + ChunkStrategy string `orm:"chunk_strategy" json:"chunk_strategy"` Status int `orm:"status" json:"status"` CreatedAt *gtime.Time `orm:"created_at" json:"created_at"` UpdatedAt *gtime.Time `orm:"updated_at" json:"updated_at"` diff --git a/kb/model/entity/model_config.go b/kb/model/entity/model_config.go index b436783..bf1b1c9 100644 --- a/kb/model/entity/model_config.go +++ b/kb/model/entity/model_config.go @@ -11,6 +11,7 @@ type ModelConfig struct { ApiKey string `orm:"api_key" json:"api_key"` Dimension int `orm:"dimension" json:"dimension"` Extra string `orm:"extra" json:"extra"` + IsDefault int `orm:"is_default" json:"is_default"` CreatedAt *gtime.Time `orm:"created_at" json:"created_at"` UpdatedAt *gtime.Time `orm:"updated_at" json:"updated_at"` } diff --git a/kb/model/entity/system_config.go b/kb/model/entity/system_config.go deleted file mode 100644 index c9596f8..0000000 --- a/kb/model/entity/system_config.go +++ /dev/null @@ -1,10 +0,0 @@ -package entity - -import "github.com/gogf/gf/v2/os/gtime" - -type SystemConfig struct { - Id int64 `orm:"id" json:"id"` - CfgKey string `orm:"cfg_key" json:"cfg_key"` - CfgValue string `orm:"cfg_value" json:"cfg_value"` - UpdatedAt *gtime.Time `orm:"updated_at" json:"updated_at"` -} diff --git a/kb/service/chat_service.go b/kb/service/chat_service.go index 2088309..aac54ec 100644 --- a/kb/service/chat_service.go +++ b/kb/service/chat_service.go @@ -152,28 +152,32 @@ func (e *OpenAIEmbedder) Dim() int { return consts.DefaultEmbeddingDim } +// EmbedStrings 内部分批请求(dashscope 单次最多 20 条),返回与 texts 同序的向量 func (e *OpenAIEmbedder) EmbedStrings(ctx context.Context, texts []string, opts ...eembedding.Option) ([][]float64, error) { - payload := map[string]any{"model": e.cfg.ModelName, "input": texts} - body, err := postOpenAI(ctx, e.cfg, strings.TrimRight(e.cfg.EndpointUrl, "/")+"/embeddings", payload) - if err != nil { - return nil, err - } - var resp struct { - Data []struct { - Embedding []float64 `json:"embedding"` - Index int `json:"index"` - } `json:"data"` - } - if err := json.Unmarshal(body, &resp); err != nil { - return nil, gerror.Wrap(err, "解析 embedding 响应失败") - } - if len(resp.Data) == 0 { - return nil, gerror.New("embedding 接口返回空数据") - } out := make([][]float64, len(texts)) - for _, d := range resp.Data { - if d.Index >= 0 && d.Index < len(out) { - out[d.Index] = d.Embedding + for start := 0; start < len(texts); start += consts.EmbedBatchSize { + end := min(start+consts.EmbedBatchSize, len(texts)) + payload := map[string]any{"model": e.cfg.ModelName, "input": texts[start:end]} + body, err := postOpenAI(ctx, e.cfg, strings.TrimRight(e.cfg.EndpointUrl, "/")+"/embeddings", payload) + if err != nil { + return nil, err + } + var resp struct { + Data []struct { + Embedding []float64 `json:"embedding"` + Index int `json:"index"` + } `json:"data"` + } + if err := json.Unmarshal(body, &resp); err != nil { + return nil, gerror.Wrap(err, "解析 embedding 响应失败") + } + if len(resp.Data) == 0 { + return nil, gerror.New("embedding 接口返回空数据") + } + for _, d := range resp.Data { + if d.Index >= 0 && d.Index < end-start { + out[start+d.Index] = d.Embedding + } } } return out, nil @@ -323,12 +327,12 @@ func (s *chatService) Ask(ctx context.Context, datasetId int64, question string, g.Log().Warningf(ctx, "graph enhance failed: %v", err) } - defaultChatModel, _, err := SystemConfigService.GetSettings(ctx) + defaultChatModel, err := dao.ModelConfig.GetDefault(ctx, consts.ModelTypeChat) if err != nil { return "", nil, err } if defaultChatModel <= 0 { - return "", nil, gerror.New("请先在设置中选择默认对话模型") + return "", nil, gerror.New("请先在设置中为对话模型设置默认") } model, err := BuildChatModel(ctx, defaultChatModel) if err != nil { @@ -451,6 +455,7 @@ func postOpenAIStream(ctx context.Context, cfg *entity.ModelConfig, url string, } func doOpenAIRequest(ctx context.Context, cfg *entity.ModelConfig, url string, payload any) (io.ReadCloser, error) { + start := time.Now() buf, err := json.Marshal(payload) if err != nil { return nil, err @@ -465,12 +470,16 @@ func doOpenAIRequest(ctx context.Context, cfg *entity.ModelConfig, url string, p } resp, err := httpClient.Do(req) if err != nil { + g.Log().Errorf(ctx, "model call failed: model=%s url=%s err=%v", cfg.ModelName, url, err) return nil, err } if resp.StatusCode >= 400 { msg, _ := io.ReadAll(resp.Body) resp.Body.Close() - return nil, gerror.Newf("模型接口 %s 返回 %d: %s", url, resp.StatusCode, strings.TrimSpace(string(msg))) + err := gerror.Newf("模型接口 %s 返回 %d: %s", url, resp.StatusCode, strings.TrimSpace(string(msg))) + g.Log().Errorf(ctx, "model call failed: model=%s url=%s status=%d err=%v", cfg.ModelName, url, resp.StatusCode, err) + return nil, err } + g.Log().Infof(ctx, "model call ok: model=%s url=%s status=%d dur=%s", cfg.ModelName, url, resp.StatusCode, time.Since(start).Round(time.Millisecond)) return resp.Body, nil } diff --git a/kb/service/chunk_service.go b/kb/service/chunk_service.go index 1546efc..8fc538d 100644 --- a/kb/service/chunk_service.go +++ b/kb/service/chunk_service.go @@ -11,17 +11,80 @@ import ( "rag-local/kb/model/domain" "rag-local/kb/model/entity" + "github.com/cloudwego/eino-ext/components/document/transformer/splitter/recursive" + "github.com/cloudwego/eino-ext/components/document/transformer/splitter/semantic" + "github.com/cloudwego/eino/components/document" eembedding "github.com/cloudwego/eino/components/embedding" + "github.com/cloudwego/eino/schema" "github.com/gogf/gf/v2/errors/gerror" "github.com/gogf/gf/v2/frame/g" ) +// 中文分块分隔符:先段落/换行,再句读标点;runeLen 保证按字计数(与标题策略一致) +var chunkSeparators = []string{"\n\n", "\n", "。", "!", "?", ";", ",", " ", ""} +var runeLen = utf8.RuneCountInString + +// SplitByStrategy 按数据集分块策略分块;semantic 需要 embedder,其余策略可传 nil +func (s *chunkService) SplitByStrategy(ctx context.Context, strategy, text string, chunkSize, overlap int, embedder eembedding.Embedder) ([]string, error) { + switch strategy { + case consts.ChunkStrategyRecursive: + sp, err := recursive.NewSplitter(ctx, &recursive.Config{ + ChunkSize: chunkSize, + OverlapSize: overlap, + Separators: chunkSeparators, + LenFunc: runeLen, + KeepType: recursive.KeepTypeEnd, + }) + if err != nil { + return nil, gerror.Wrap(err, "构建递归分块器失败") + } + return s.transformText(ctx, sp, text) + case consts.ChunkStrategySemantic: + if embedder == nil { + return nil, gerror.New("语义分块需要向量模型") + } + sp, err := semantic.NewSplitter(ctx, &semantic.Config{ + Embedding: embedder, + BufferSize: 1, + MinChunkSize: chunkSize / 2, + Separators: chunkSeparators[:len(chunkSeparators)-2], + LenFunc: runeLen, + }) + if err != nil { + return nil, gerror.Wrap(err, "构建语义分块器失败") + } + return s.transformText(ctx, sp, text) + default: + return s.SplitText(text, chunkSize, overlap), nil + } +} + +func (s *chunkService) transformText(ctx context.Context, sp document.Transformer, text string) ([]string, error) { + docs, err := sp.Transform(ctx, []*schema.Document{{ID: "0", Content: text}}) + if err != nil { + return nil, gerror.Wrap(err, "分块失败") + } + out := make([]string, 0, len(docs)) + for _, d := range docs { + if c := strings.TrimSpace(d.Content); c != "" { + out = append(out, c) + } + } + return out, nil +} + var ChunkService = &chunkService{} type chunkService struct{} // SplitText 文本分块:标题感知 + 固定大小回退,超长段落按句号/换行切分并保留重叠 -func (s *chunkService) SplitText(text string) []string { +func (s *chunkService) SplitText(text string, chunkSize, overlap int) []string { + if chunkSize <= 0 { + chunkSize = consts.DefaultChunkSize + } + if overlap < 0 { + overlap = consts.DefaultChunkOverlap + } text = strings.ReplaceAll(text, "\r\n", "\n") text = strings.ReplaceAll(text, "\r", "\n") @@ -31,7 +94,7 @@ func (s *chunkService) SplitText(text string) []string { curRunes := 0 for _, p := range paragraphs { pRunes := utf8.RuneCountInString(p) - if cur.Len() > 0 && curRunes+pRunes > consts.MaxChunkSize { + if cur.Len() > 0 && curRunes+pRunes > chunkSize { merged = append(merged, cur.String()) cur.Reset() curRunes = 0 @@ -46,11 +109,11 @@ func (s *chunkService) SplitText(text string) []string { var chunks []string for _, c := range merged { - if len(c) <= consts.MaxChunkSize { + if len(c) <= chunkSize { chunks = append(chunks, strings.TrimSpace(c)) continue } - chunks = append(chunks, forceSplit(c)...) + chunks = append(chunks, forceSplit(c, chunkSize, overlap)...) } return chunks } @@ -87,19 +150,19 @@ func isHeading(line string) bool { } // forceSplit 超长文本按可读位置切分(rune 安全,避免切在 UTF-8 中间),重叠 overlap 字 -func forceSplit(text string) []string { +func forceSplit(text string, chunkSize, overlap int) []string { runes := []rune(text) var result []string - for len(runes) > consts.MaxChunkSize { - limit := min(consts.MaxChunkSize, len(runes)) + for len(runes) > chunkSize { + limit := min(chunkSize, len(runes)) cut := lastCutPoint(runes[:limit]) - if cut < consts.MaxChunkSize/2 { - cut = consts.MaxChunkSize + if cut < chunkSize/2 { + cut = chunkSize } if chunk := strings.TrimSpace(string(runes[:cut])); chunk != "" { result = append(result, chunk) } - runes = runes[max(0, cut-consts.ChunkOverlap):] + runes = runes[max(0, cut-overlap):] } if chunk := strings.TrimSpace(string(runes)); chunk != "" { result = append(result, chunk) diff --git a/kb/service/dataset_service.go b/kb/service/dataset_service.go index 0d4d56a..8252213 100644 --- a/kb/service/dataset_service.go +++ b/kb/service/dataset_service.go @@ -23,15 +23,70 @@ func (s *datasetService) Save(ctx context.Context, m *entity.Dataset) (int64, er if m.Status == 0 { m.Status = 1 } + if err := s.validateStrategy(m); err != nil { + return 0, err + } if m.Id > 0 { + old, err := dao.Dataset.GetOne(ctx, m.Id) + if err != nil { + return 0, err + } if err := dao.Dataset.Update(ctx, m); err != nil { return 0, err } + // 向量模型变更 → 入队重新向量化任务,保证 chunk 向量与提问检索模型一致 + if old != nil && old.EmbeddingCfgId != m.EmbeddingCfgId { + if err := s.enqueueTask(ctx, m.Id, consts.TaskTypeReembed); err != nil { + g.Log().Warningf(ctx, "enqueue reembed for dataset %d failed: %v", m.Id, err) + } + } + // 分块策略/大小/重叠变更 → 入队完整重新解析任务(重新分词+向量化) + if old != nil && (old.ChunkSize != m.ChunkSize || old.ChunkOverlap != m.ChunkOverlap || old.ChunkStrategy != m.ChunkStrategy) { + if err := s.enqueueTask(ctx, m.Id, consts.TaskTypeParse); err != nil { + g.Log().Warningf(ctx, "enqueue reparse for dataset %d failed: %v", m.Id, err) + } + } return m.Id, nil } + if m.ChunkSize <= 0 { + m.ChunkSize = consts.DefaultChunkSize + } + if m.ChunkOverlap < 0 { + m.ChunkOverlap = consts.DefaultChunkOverlap + } + if m.ChunkStrategy == "" { + m.ChunkStrategy = consts.ChunkStrategyTitle + } return dao.Dataset.Insert(ctx, m) } +func (s *datasetService) validateStrategy(m *entity.Dataset) error { + if m.ChunkStrategy == "" { + m.ChunkStrategy = consts.ChunkStrategyTitle + } + if m.EmbeddingCfgId == 0 { + return gerror.New("数据集必须绑定向量模型,请先选择向量模型") + } + // 语义分块无重叠参数,清零避免误导 + if m.ChunkStrategy == consts.ChunkStrategySemantic { + m.ChunkOverlap = 0 + } + return nil +} + +func (s *datasetService) enqueueTask(ctx context.Context, datasetId int64, taskType string) error { + docs, _, err := dao.Document.List(ctx, datasetId, 1, 100000) + if err != nil { + return err + } + for _, doc := range docs { + if _, err := dao.ParseTask.Insert(ctx, doc.Id, datasetId, taskType); err != nil { + return err + } + } + return nil +} + func (s *datasetService) Delete(ctx context.Context, id int64) error { // 有文档的数据集不允许删除 count, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameDocument).Ctx(ctx). diff --git a/kb/service/kg_entity_service.go b/kb/service/kg_entity_service.go index 71b42a8..ebcbe2d 100644 --- a/kb/service/kg_entity_service.go +++ b/kb/service/kg_entity_service.go @@ -5,6 +5,7 @@ import ( "encoding/json" "strings" + "rag-local/kb/consts" "rag-local/kb/dao" "rag-local/kb/model/entity" @@ -54,12 +55,12 @@ func (s *kgEntityService) ExtractDocument(ctx context.Context, datasetId, docume } func (s *kgEntityService) buildModel(ctx context.Context) (*OpenAIChatModel, error) { - defaultChatModel, _, err := SystemConfigService.GetSettings(ctx) + defaultChatModel, err := dao.ModelConfig.GetDefault(ctx, consts.ModelTypeChat) if err != nil { return nil, err } if defaultChatModel <= 0 { - return nil, gerror.New("未配置默认对话模型") + return nil, gerror.New("未设置默认对话模型") } return BuildChatModel(ctx, defaultChatModel) } diff --git a/kb/service/model_config_service.go b/kb/service/model_config_service.go index d859b69..c36a030 100644 --- a/kb/service/model_config_service.go +++ b/kb/service/model_config_service.go @@ -42,6 +42,18 @@ func (s *modelConfigService) Save(ctx context.Context, m *entity.ModelConfig) (i 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) diff --git a/kb/service/parse_task_service.go b/kb/service/parse_task_service.go index f246a99..4a0bfe8 100644 --- a/kb/service/parse_task_service.go +++ b/kb/service/parse_task_service.go @@ -10,7 +10,6 @@ import ( "rag-local/kb/dao" "rag-local/kb/model/entity" - eembedding "github.com/cloudwego/eino/components/embedding" "github.com/gogf/gf/v2/errors/gerror" "github.com/gogf/gf/v2/frame/g" ) @@ -58,6 +57,10 @@ func (s *parseTaskService) processOne(ctx context.Context) { s.fail(ctx, task, "文档不存在") return } + if task.TaskType == consts.TaskTypeReembed { + s.processReembed(ctx, task) + return + } if err := dao.Document.UpdateFields(ctx, doc.Id, g.Map{"status": consts.DocumentStatusParsing}); err != nil { s.fail(ctx, task, "更新文档状态失败: "+err.Error()) return @@ -68,18 +71,38 @@ func (s *parseTaskService) processOne(ctx context.Context) { s.fail(ctx, task, "解析文件失败: "+err.Error()) return } - chunks := ChunkService.SplitText(text) - // 数据集绑定 embedding 配置时构建向量模型,无配置降级为仅全文索引 - var embedder eembedding.Embedder - if cfgId, err := dao.Dataset.GetEmbeddingCfgId(ctx, task.DatasetId); err == nil && cfgId > 0 { - if em, err := BuildEmbedder(ctx, cfgId); err == nil { - embedder = em - if dim := em.Dim(); dim != g.Cfg().MustGet(ctx, "vector.dim", consts.DefaultEmbeddingDim).Int() { - g.Log().Warningf(ctx, "embedding 维度 %d 与 vec0 表维度不一致,请确认 vector.dim 配置", dim) - } - } else { - g.Log().Warningf(ctx, "build embedder failed, fallback to fts only: %v", err) - } + // 数据集分块配置(策略/大小/重叠),未设置用默认值 + chunkSize, chunkOverlap, strategy := consts.DefaultChunkSize, consts.DefaultChunkOverlap, consts.ChunkStrategyTitle + if ds, err := dao.Dataset.GetOne(ctx, task.DatasetId); err == nil && ds != nil { + chunkSize, chunkOverlap, strategy = ds.ChunkSize, ds.ChunkOverlap, ds.ChunkStrategy + } + // 数据集必须绑定向量模型,构建失败直接失败任务(不降级全文索引) + cfgId, err := dao.Dataset.GetEmbeddingCfgId(ctx, task.DatasetId) + if err != nil { + s.fail(ctx, task, "读取数据集向量模型配置失败: "+err.Error()) + return + } + if cfgId <= 0 { + s.fail(ctx, task, "数据集未绑定向量模型") + return + } + embedder, err := BuildEmbedder(ctx, cfgId) + if err != nil { + s.fail(ctx, task, "构建向量模型失败: "+err.Error()) + return + } + if dim := embedder.Dim(); dim != g.Cfg().MustGet(ctx, "vector.dim", consts.DefaultEmbeddingDim).Int() { + g.Log().Warningf(ctx, "embedding 维度 %d 与 vec0 表维度不一致,请确认 vector.dim 配置", dim) + } + // 重新解析时先清空旧分块,避免新旧分块并存 + if err := ChunkService.DeleteByDocument(ctx, doc.Id); err != nil { + s.fail(ctx, task, "清理旧分块失败: "+err.Error()) + return + } + chunks, err := ChunkService.SplitByStrategy(ctx, strategy, text, chunkSize, chunkOverlap, embedder) + if err != nil { + s.fail(ctx, task, "分块失败: "+err.Error()) + return } if err := ChunkService.InsertAll(ctx, task.DatasetId, doc.Id, chunks, embedder); err != nil { s.fail(ctx, task, "写入分块失败: "+err.Error()) @@ -94,6 +117,18 @@ func (s *parseTaskService) processOne(ctx context.Context) { } } +// processReembed 重新向量化任务:分块文本不变,用数据集当前绑定模型重算全部向量;失败不动文档状态 +func (s *parseTaskService) processReembed(ctx context.Context, task *entity.ParseTask) { + if err := DocumentService.Reembed(ctx, task.DocumentId); err != nil { + _ = dao.ParseTask.UpdateStatus(ctx, task.Id, consts.TaskStatusFailed, "重新向量化失败: "+err.Error()) + g.Log().Errorf(ctx, "reembed task %d failed: %s", task.Id, err.Error()) + return + } + if err := dao.ParseTask.UpdateStatus(ctx, task.Id, consts.TaskStatusDone, ""); err != nil { + g.Log().Errorf(ctx, "mark task done failed: %v", err) + } +} + func (s *parseTaskService) fail(ctx context.Context, task *entity.ParseTask, msg string) { _ = dao.ParseTask.UpdateStatus(ctx, task.Id, consts.TaskStatusFailed, msg) _ = dao.Document.UpdateFields(ctx, task.DocumentId, g.Map{ diff --git a/kb/service/system_config_service.go b/kb/service/system_config_service.go index 81bcb7b..b9e3dcc 100644 --- a/kb/service/system_config_service.go +++ b/kb/service/system_config_service.go @@ -2,99 +2,26 @@ package service import ( "context" - "time" "rag-local/common" - "rag-local/kb/consts" - "rag-local/kb/dao" "github.com/gogf/gf/v2/errors/gerror" - "github.com/gogf/gf/v2/os/gcache" - "github.com/gogf/gf/v2/util/gconv" ) var SystemConfigService = &systemConfigService{} type systemConfigService struct{} -func init() { - // 注入指纹校验函数,避免 common → service 循环依赖 - common.CheckTokenFingerprint = SystemConfigService.CheckTokenFingerprint -} - -// EnsureAccessToken 首次启动生成访问令牌并写入配置;已有则复用。返回当前令牌 +// EnsureAccessToken 每次启动生成新的访问令牌(内存持有,不落库),供登录使用 func (s *systemConfigService) EnsureAccessToken(ctx context.Context) (string, error) { - token, err := dao.SystemConfig.Get(ctx, consts.CfgKeyAccessToken) - if err != nil { - return "", err - } - if token == "" { - token = common.RandomToken(16) - if err := dao.SystemConfig.Set(ctx, consts.CfgKeyAccessToken, token); err != nil { - return "", err - } - } + token := common.RandomToken(16) + common.SetAccessToken(token) return token, nil } func (s *systemConfigService) Login(ctx context.Context, token string) (string, error) { - cur, err := dao.SystemConfig.Get(ctx, consts.CfgKeyAccessToken) - if err != nil { - return "", err - } - if cur == "" || token != cur { + if !common.CheckAccessToken(token) { return "", gerror.New("访问令牌错误") } - return common.SignToken("owner", common.TokenFingerprint(cur), common.TokenExpireSeconds) -} - -// GetToken 当前访问令牌(设置页展示用) -func (s *systemConfigService) GetToken(ctx context.Context) (string, error) { - return dao.SystemConfig.Get(ctx, consts.CfgKeyAccessToken) -} - -// TokenFingerprint 当前令牌指纹(带短缓存,鉴权路径避免频繁查库) -func (s *systemConfigService) TokenFingerprint(ctx context.Context) string { - if v, err := gcache.Get(ctx, "access_token_fp"); err == nil && !v.IsNil() { - return v.String() - } - token, err := dao.SystemConfig.Get(ctx, consts.CfgKeyAccessToken) - if err != nil || token == "" { - return "" - } - fp := common.TokenFingerprint(token) - _ = gcache.Set(ctx, "access_token_fp", fp, 10*time.Second) - return fp -} - -func (s *systemConfigService) CheckTokenFingerprint(ctx context.Context, fp string) bool { - return fp != "" && fp == s.TokenFingerprint(ctx) -} - -func (s *systemConfigService) RegenerateToken(ctx context.Context) (string, error) { - token := common.RandomToken(16) - if err := dao.SystemConfig.Set(ctx, consts.CfgKeyAccessToken, token); err != nil { - return "", err - } - _, _ = gcache.Remove(ctx, "access_token_fp") - return token, nil -} - -func (s *systemConfigService) GetSettings(ctx context.Context) (defaultChatModel, defaultDataset int64, err error) { - v1, err := dao.SystemConfig.Get(ctx, consts.CfgKeyDefaultChatModel) - if err != nil { - return - } - v2, err := dao.SystemConfig.Get(ctx, consts.CfgKeyDefaultDataset) - if err != nil { - return - } - return gconv.Int64(v1), gconv.Int64(v2), nil -} - -func (s *systemConfigService) UpdateSettings(ctx context.Context, defaultChatModel, defaultDataset int64) error { - if err := dao.SystemConfig.Set(ctx, consts.CfgKeyDefaultChatModel, gconv.String(defaultChatModel)); err != nil { - return err - } - return dao.SystemConfig.Set(ctx, consts.CfgKeyDefaultDataset, gconv.String(defaultDataset)) + return common.SignToken("owner", common.AccessTokenFingerprint(), common.TokenExpireSeconds) } diff --git a/ui-src/src/api/auth.js b/ui-src/src/api/auth.js index 994c4f7..8dd816b 100644 --- a/ui-src/src/api/auth.js +++ b/ui-src/src/api/auth.js @@ -3,19 +3,3 @@ import request from './request.js' export function login(data) { return request.post('/system-config/login', data) } - -export function getSystemConfig() { - return request.get('/system-config') -} - -export function updateSystemConfig(data) { - return request.put('/system-config', data) -} - -export function regenerateToken() { - return request.post('/system-config/regenerate-token') -} - -export function getToken() { - return request.get('/system-config/token') -} diff --git a/ui-src/src/api/model_config.js b/ui-src/src/api/model_config.js index cbfd038..6e8c21c 100644 --- a/ui-src/src/api/model_config.js +++ b/ui-src/src/api/model_config.js @@ -15,3 +15,7 @@ export function deleteModelConfig(id) { export function testModelConfig(id) { return request.post('/model-config/test', { id }) } + +export function setDefaultModelConfig(id) { + return request.post('/model-config/set-default', { id }) +} diff --git a/ui-src/src/views/Chat.vue b/ui-src/src/views/Chat.vue index d02fad6..583465d 100644 --- a/ui-src/src/views/Chat.vue +++ b/ui-src/src/views/Chat.vue @@ -89,7 +89,8 @@ function scrollBottom() { } async function refreshConversations() { - conversations.value = await listConversations() + const c = await listConversations() + conversations.value = (c && c.list) || [] } function newConversation() { @@ -100,7 +101,8 @@ function newConversation() { async function openConversation(c) { currentConvId.value = c.id if (c.dataset_id) datasetId.value = c.dataset_id - const list = await listMessages(c.id) + const r = await listMessages(c.id) + const list = (r && r.list) || [] messages.value = list.map(m => { let citations = [] if (m.citations) { diff --git a/ui-src/src/views/DatasetDetail.vue b/ui-src/src/views/DatasetDetail.vue index 5c965a7..7e51d4b 100644 --- a/ui-src/src/views/DatasetDetail.vue +++ b/ui-src/src/views/DatasetDetail.vue @@ -62,7 +62,7 @@ -
保存后重新分词并向量化,向量模型未绑定时仅更新全文索引
+
保存后重新分词并向量化