diff --git a/data/business.db b/data/business.db index 98c8824..930dc2e 100644 Binary files a/data/business.db and b/data/business.db differ diff --git a/data/chat.db b/data/chat.db index 480946f..30c4cc0 100644 Binary files a/data/chat.db and b/data/chat.db differ diff --git a/kb/consts/consts.go b/kb/consts/consts.go index a8434d7..e188c9d 100644 --- a/kb/consts/consts.go +++ b/kb/consts/consts.go @@ -8,17 +8,28 @@ const ( HybridTopK = 5 // 混合检索最终返回数(重排+门槛过滤后) RrfK = 60 // RRF 融合常数 - RerankTopK = 10 // 喂给 LLM 重排器的候选数(RRF 融合后截取) + RerankTopK = 10 // 喂给 LLM 重排器的候选数(RRF 融合后截取;数据集可覆盖) RerankKeepRatio = 0.5 // 重排分低于最高分该比例的条目剔除(门槛作用在语义分上,RRF 分因区间过窄无区分度) RerankMinScore = 6 // 重排分绝对下限:低于该分的条目直接剔除(重排器对泛化条款会给宽松低分,需绝对下限兜底) RerankMaxChars = 500 // 重排候选段截断字数(控制 prompt 长度) + // 数据集召回数量配置(vec_top_k/fts_top_k/rerank_top_k/recall_top_k):0=用上方全局默认,-1=尽量多,>0=固定值 + MaxRecallTopK = 50 // 四个配置字段的统一上限 + MaxRecallRawTopK = 100 // 原始召回(向量/全文)"-1=尽量多"的保护上限 + MaxRerankTopK = 60 // 重排候选 "-1=尽量多"的保护上限(重排 prompt 成本护栏) + MaxRecallChars = 50000 // 最终返回 "-1=不限制数量"时的总字符物理保护(按相关性门槛过滤后仍超预算则按分截断) + DefaultChunkSize = 800 // 分块最大字数(数据集默认值) DefaultChunkOverlap = 150 // 分块重叠字数(数据集默认值) + // 智能体(ReAct)参数 + MaxReactRounds = 10 // ReAct 轮次上限 + ToolResultMaxChars = 1500 // search 工具返回的单条 chunk 截断字数(控制上下文体积) + // 全局设置键(app_config 表) SettingsKeyChunkSize = "chunk_default_size" SettingsKeyChunkOverlap = "chunk_default_overlap" + SettingsKeyReactRounds = "react_default_rounds" ParsePollIntervalSeconds = 5 // 解析/标注任务轮询间隔 diff --git a/kb/controller/dataset_controller.go b/kb/controller/dataset_controller.go index 058dcd4..4a87742 100644 --- a/kb/controller/dataset_controller.go +++ b/kb/controller/dataset_controller.go @@ -28,6 +28,11 @@ func (c *dataset) Save(ctx context.Context, req *dto.SaveDatasetReq) (*dto.SaveD EmbeddingCfgId: req.EmbeddingCfgId, ChunkSize: req.ChunkSize, ChunkOverlap: req.ChunkOverlap, + ReactRounds: req.ReactRounds, + VecTopK: req.VecTopK, + FtsTopK: req.FtsTopK, + RerankTopK: req.RerankTopK, + RecallTopK: req.RecallTopK, Status: 1, }) if err != nil { diff --git a/kb/controller/message_controller.go b/kb/controller/message_controller.go index 4ce09cb..0163647 100644 --- a/kb/controller/message_controller.go +++ b/kb/controller/message_controller.go @@ -68,6 +68,9 @@ func (c *message) Chat(ctx context.Context, req *dto.ChatReq) (*dto.ChatRes, err }, func(delta string) { send("delta", map[string]string{"content": delta}) + }, + func(thinking string) { + send("thinking", map[string]string{"type": "thinking", "message": thinking}) }) if err != nil { send("error", map[string]string{"message": err.Error()}) diff --git a/kb/controller/system_config_controller.go b/kb/controller/system_config_controller.go index 1b8495b..178d18d 100644 --- a/kb/controller/system_config_controller.go +++ b/kb/controller/system_config_controller.go @@ -20,15 +20,15 @@ func (c *systemConfig) Login(ctx context.Context, req *dto.LoginReq) (res *dto.L } func (c *systemConfig) GetSettings(ctx context.Context, _ *dto.GetSettingsReq) (*dto.GetSettingsRes, error) { - size, overlap, err := service.SystemConfigService.GetSettings(ctx) + size, overlap, rounds, err := service.SystemConfigService.GetSettings(ctx) if err != nil { return nil, err } - return &dto.GetSettingsRes{ChunkSize: size, ChunkOverlap: overlap}, nil + return &dto.GetSettingsRes{ChunkSize: size, ChunkOverlap: overlap, ReactRounds: rounds}, nil } func (c *systemConfig) SaveSettings(ctx context.Context, req *dto.SaveSettingsReq) (*dto.SaveSettingsRes, error) { - if err := service.SystemConfigService.SaveSettings(ctx, req.ChunkSize, req.ChunkOverlap); err != nil { + if err := service.SystemConfigService.SaveSettings(ctx, req.ChunkSize, req.ChunkOverlap, req.ReactRounds); err != nil { return nil, err } return &dto.SaveSettingsRes{}, nil diff --git a/kb/dao/dataset_dao.go b/kb/dao/dataset_dao.go index c8e6e16..82363c6 100644 --- a/kb/dao/dataset_dao.go +++ b/kb/dao/dataset_dao.go @@ -45,6 +45,11 @@ func init() { {"chunk_strategy", "chunk_strategy TEXT NOT NULL DEFAULT 'title'"}, {"unit_pattern", "unit_pattern TEXT NOT NULL DEFAULT ''"}, {"context_pattern", "context_pattern TEXT NOT NULL DEFAULT ''"}, + {"react_rounds", "react_rounds INTEGER NOT NULL DEFAULT 0"}, + {"vec_top_k", "vec_top_k INTEGER NOT NULL DEFAULT 0"}, + {"fts_top_k", "fts_top_k INTEGER NOT NULL DEFAULT 0"}, + {"rerank_top_k", "rerank_top_k INTEGER NOT NULL DEFAULT 0"}, + {"recall_top_k", "recall_top_k INTEGER NOT NULL DEFAULT 0"}, } { cnt, err := g.DB(consts.DbGroupDefault).Ctx(ctx).GetValue(ctx, "SELECT COUNT(*) FROM pragma_table_info('"+consts.TableNameDataset+"') WHERE name=?", col.name) @@ -87,6 +92,11 @@ func (d *datasetDao) Insert(ctx context.Context, data *entity.Dataset) (int64, e "embedding_cfg_id": data.EmbeddingCfgId, "chunk_size": data.ChunkSize, "chunk_overlap": data.ChunkOverlap, + "react_rounds": data.ReactRounds, + "vec_top_k": data.VecTopK, + "fts_top_k": data.FtsTopK, + "rerank_top_k": data.RerankTopK, + "recall_top_k": data.RecallTopK, "status": data.Status, "created_at": now, "updated_at": now, @@ -104,6 +114,11 @@ func (d *datasetDao) Update(ctx context.Context, data *entity.Dataset) error { "embedding_cfg_id": data.EmbeddingCfgId, "chunk_size": data.ChunkSize, "chunk_overlap": data.ChunkOverlap, + "react_rounds": data.ReactRounds, + "vec_top_k": data.VecTopK, + "fts_top_k": data.FtsTopK, + "rerank_top_k": data.RerankTopK, + "recall_top_k": data.RecallTopK, "status": data.Status, "updated_at": gtime.Now().Format("Y-m-d H:i:s"), }).Where("id", data.Id).Update() @@ -132,3 +147,42 @@ func (d *datasetDao) GetEmbeddingCfgId(ctx context.Context, id int64) (int64, er } return r["embedding_cfg_id"].Int64(), nil } + +// GetReactRounds 读取数据集 ReAct 轮次配置(未配置/异常时返回 0 = 关闭智能体模式) +func (d *datasetDao) GetReactRounds(ctx context.Context, id int64) (int, error) { + r, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameDataset).Ctx(ctx). + Fields("react_rounds").Where("id", id).One() + if err != nil { + return 0, err + } + if r == nil { + return 0, nil + } + return r["react_rounds"].Int(), nil +} + +// RecallParams 数据集召回数量配置(0=全局默认,-1=尽量多,>0=固定值) +type RecallParams struct { + VecTopK int + FtsTopK int + RerankTopK int + RecallTopK int +} + +// GetRecallParams 读取数据集召回数量配置(未配置/异常时返回全 0 = 用全局默认) +func (d *datasetDao) GetRecallParams(ctx context.Context, id int64) (*RecallParams, error) { + r, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameDataset).Ctx(ctx). + Fields("vec_top_k,fts_top_k,rerank_top_k,recall_top_k").Where("id", id).One() + if err != nil { + return nil, err + } + if r == nil { + return &RecallParams{}, nil + } + return &RecallParams{ + VecTopK: r["vec_top_k"].Int(), + FtsTopK: r["fts_top_k"].Int(), + RerankTopK: r["rerank_top_k"].Int(), + RecallTopK: r["recall_top_k"].Int(), + }, nil +} diff --git a/kb/model/domain/retrieval.go b/kb/model/domain/retrieval.go index bbee6bb..982c877 100644 --- a/kb/model/domain/retrieval.go +++ b/kb/model/domain/retrieval.go @@ -40,6 +40,9 @@ type Citation struct { Content string `json:"content"` Score float64 `json:"score"` Sources []string `json:"sources"` + // HighlightStart/End 内容中最相关段落的 UTF-16 偏移(0=未命中),供前端高亮标注 + HighlightStart int `json:"highlight_start,omitempty"` + HighlightEnd int `json:"highlight_end,omitempty"` } // VecJson 向量 JSON 序列化([0.1,0.2,...]) diff --git a/kb/model/dto/dataset_dto.go b/kb/model/dto/dataset_dto.go index 9bc7dd8..3bd20d8 100644 --- a/kb/model/dto/dataset_dto.go +++ b/kb/model/dto/dataset_dto.go @@ -22,6 +22,11 @@ type SaveDatasetReq struct { EmbeddingCfgId int64 `json:"embedding_cfg_id"` ChunkSize int `json:"chunk_size"` ChunkOverlap int `json:"chunk_overlap"` + ReactRounds int `json:"react_rounds"` + VecTopK int `json:"vec_top_k"` + FtsTopK int `json:"fts_top_k"` + RerankTopK int `json:"rerank_top_k"` + RecallTopK int `json:"recall_top_k"` } type SaveDatasetRes struct { diff --git a/kb/model/dto/system_config_dto.go b/kb/model/dto/system_config_dto.go index 00b084b..c7fc049 100644 --- a/kb/model/dto/system_config_dto.go +++ b/kb/model/dto/system_config_dto.go @@ -18,12 +18,14 @@ type GetSettingsReq struct { type GetSettingsRes struct { ChunkSize int `json:"chunk_size"` // 默认分块大小 ChunkOverlap int `json:"chunk_overlap"` // 默认重叠字数 + ReactRounds int `json:"react_rounds"` // 默认智能体轮次(0=关闭) } type SaveSettingsReq struct { g.Meta `path:"/save-settings" method:"post" tags:"系统配置" summary:"保存全局设置"` ChunkSize int `json:"chunk_size"` ChunkOverlap int `json:"chunk_overlap"` + ReactRounds int `json:"react_rounds"` } type SaveSettingsRes struct{} diff --git a/kb/model/entity/dataset.go b/kb/model/entity/dataset.go index 3f60ba8..6ac96dd 100644 --- a/kb/model/entity/dataset.go +++ b/kb/model/entity/dataset.go @@ -3,15 +3,21 @@ package entity import "github.com/gogf/gf/v2/os/gtime" type Dataset struct { - Id int64 `orm:"id" json:"id"` - 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"` - UnitPattern string `orm:"unit_pattern" json:"unit_pattern"` - ContextPattern string `orm:"context_pattern" json:"context_pattern"` - Status int `orm:"status" json:"status"` - CreatedAt *gtime.Time `orm:"created_at" json:"created_at"` - UpdatedAt *gtime.Time `orm:"updated_at" json:"updated_at"` + Id int64 `orm:"id" json:"id"` + 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"` + UnitPattern string `orm:"unit_pattern" json:"unit_pattern"` + ContextPattern string `orm:"context_pattern" json:"context_pattern"` + ReactRounds int `orm:"react_rounds" json:"react_rounds"` + // 召回数量配置(0=全局默认,-1=尽量多,>0=固定值) + VecTopK int `orm:"vec_top_k" json:"vec_top_k"` // 向量原始召回数 + FtsTopK int `orm:"fts_top_k" json:"fts_top_k"` // 全文原始召回数 + RerankTopK int `orm:"rerank_top_k" json:"rerank_top_k"` // 重排候选数 + RecallTopK int `orm:"recall_top_k" json:"recall_top_k"` // 最终返回数 + 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/service/chat_service.go b/kb/service/chat_service.go index aba6bb2..85773b2 100644 --- a/kb/service/chat_service.go +++ b/kb/service/chat_service.go @@ -10,11 +10,14 @@ import ( "io" "math" "net/http" + "regexp" "sort" "strconv" "strings" "sync" "time" + "unicode/utf16" + "unicode/utf8" "rag-local/common" "rag-local/kb/consts" @@ -35,8 +38,23 @@ var httpClient = &http.Client{Timeout: 2 * time.Minute} // ---------- OpenAI 兼容 HTTP 组件 ---------- type openAIMessage struct { - Role string `json:"role"` - Content string `json:"content"` + Role string `json:"role"` + Content string `json:"content"` + ToolCalls []openAIToolCall `json:"tool_calls,omitempty"` + ToolCallID string `json:"tool_call_id,omitempty"` +} + +// openAIToolCall OpenAI 格式工具调用(镜像 eino schema.ToolCall;Index 仅流式累积时使用) +type openAIToolCall struct { + Index *int `json:"index,omitempty"` + ID string `json:"id,omitempty"` + Type string `json:"type,omitempty"` + Function openAIToolFunction `json:"function"` +} + +type openAIToolFunction struct { + Name string `json:"name"` + Arguments string `json:"arguments"` } type openAIChatResponse struct { @@ -56,7 +74,8 @@ type openAIStreamChunk struct { // OpenAIChatModel 基于 OpenAI 兼容 /chat/completions 接口的对话模型,实现 eino model.ChatModel type OpenAIChatModel struct { - cfg *entity.ModelConfig + cfg *entity.ModelConfig + tools []*schema.ToolInfo } func NewOpenAIChatModel(cfg *entity.ModelConfig) *OpenAIChatModel { @@ -69,6 +88,7 @@ func (m *OpenAIChatModel) Generate(ctx context.Context, input []*schema.Message, "messages": buildOpenAIMessages(input), "stream": false, } + m.withTools(payload) body, err := postOpenAI(ctx, m.cfg, m.endpoint("/chat/completions"), payload) if err != nil { return nil, err @@ -81,7 +101,7 @@ func (m *OpenAIChatModel) Generate(ctx context.Context, input []*schema.Message, return nil, gerror.New("模型返回空响应") } choice := resp.Choices[0] - msg := &schema.Message{Role: schema.Assistant, Content: choice.Message.Content} + msg := &schema.Message{Role: schema.Assistant, Content: choice.Message.Content, ToolCalls: toSchemaToolCalls(choice.Message.ToolCalls)} if choice.FinishReason != "" { msg.ResponseMeta = &schema.ResponseMeta{FinishReason: choice.FinishReason, Usage: resp.Usage} } @@ -94,6 +114,7 @@ func (m *OpenAIChatModel) Stream(ctx context.Context, input []*schema.Message, o "messages": buildOpenAIMessages(input), "stream": true, } + m.withTools(payload) reader, writer := schema.Pipe[*schema.Message](16) go func() { defer writer.Close() @@ -103,6 +124,13 @@ func (m *OpenAIChatModel) Stream(ctx context.Context, input []*schema.Message, o return } defer body.Close() + // 流式 tool_call 累积:按 delta.index 分组(部分提供商省略 index 时按到达顺序), + // 流结束后若存在 tool_call 则补发一条携带完整 ToolCalls 的消息 + type tcAcc struct { + id, name, args string + } + accs := map[int]*tcAcc{} + nextIdx := 0 br := bufio.NewReader(body) for { line, err := br.ReadBytes('\n') @@ -114,25 +142,107 @@ func (m *OpenAIChatModel) Stream(ctx context.Context, input []*schema.Message, o } var chunk openAIStreamChunk if json.Unmarshal([]byte(data), &chunk) == nil && len(chunk.Choices) > 0 { - if delta := chunk.Choices[0].Delta.Content; delta != "" { - if closed := writer.Send(&schema.Message{Role: schema.Assistant, Content: delta}, nil); closed { + delta := chunk.Choices[0].Delta + if delta.Content != "" { + if closed := writer.Send(&schema.Message{Role: schema.Assistant, Content: delta.Content}, nil); closed { return } } + for _, tc := range delta.ToolCalls { + idx := nextIdx + if tc.Index != nil { + idx = *tc.Index + } else { + nextIdx++ + } + a, ok := accs[idx] + if !ok { + a = &tcAcc{} + accs[idx] = a + } + if tc.ID != "" { + a.id = tc.ID + } + if tc.Function.Name != "" { + a.name = tc.Function.Name + } + a.args += tc.Function.Arguments + } } } if err != nil { break } } + if len(accs) > 0 { + keys := make([]int, 0, len(accs)) + for k := range accs { + keys = append(keys, k) + } + sort.Ints(keys) + toolCalls := make([]schema.ToolCall, 0, len(keys)) + for _, k := range keys { + a := accs[k] + toolCalls = append(toolCalls, schema.ToolCall{ + ID: a.id, + Type: "function", + Function: schema.FunctionCall{Name: a.name, Arguments: a.args}, + }) + } + if closed := writer.Send(&schema.Message{Role: schema.Assistant, ToolCalls: toolCalls}, nil); closed { + return + } + } }() return reader, nil } func (m *OpenAIChatModel) BindTools(tools []*schema.ToolInfo) error { + m.tools = tools return nil } +// withTools 将已绑定的工具按 OpenAI function calling 格式写入请求体 +func (m *OpenAIChatModel) withTools(payload map[string]any) { + if len(m.tools) == 0 { + return + } + tools := make([]map[string]any, 0, len(m.tools)) + for _, ti := range m.tools { + if ti == nil { + continue + } + fn := map[string]any{"name": ti.Name, "description": ti.Desc} + if js, err := ti.ToJSONSchema(); err == nil && js != nil { + fn["parameters"] = js + } + tools = append(tools, map[string]any{"type": "function", "function": fn}) + } + if len(tools) > 0 { + payload["tools"] = tools + } +} + +// toSchemaToolCalls OpenAI tool_calls → eino schema.ToolCall(去除 Index,Type 缺省 function) +func toSchemaToolCalls(calls []openAIToolCall) []schema.ToolCall { + if len(calls) == 0 { + return nil + } + out := make([]schema.ToolCall, 0, len(calls)) + for _, c := range calls { + typ := c.Type + if typ == "" { + typ = "function" + } + out = append(out, schema.ToolCall{ + ID: c.ID, + Type: typ, + Function: schema.FunctionCall{Name: c.Function.Name, Arguments: c.Function.Arguments}, + }) + } + return out +} + func (m *OpenAIChatModel) endpoint(path string) string { return strings.TrimRight(m.cfg.EndpointUrl, "/") + path } @@ -226,15 +336,66 @@ type HybridRetriever struct { embedder eembedding.Embedder datasetId int64 reranker *OpenAIChatModel // LLM 重排器(默认对话模型),为 nil 时跳过重排走 RRF 顺序 + // 各阶段召回数(构建时读取数据集配置) + vecTopK int + ftsTopK int + rerankTopK int + finalTopK int + // finalUnlimited 数据集 recall_top_k=-1:最终返回不限制条数,重排后仅按相关性门槛过滤(总字符受 MaxRecallChars 保护) + finalUnlimited bool } -func NewHybridRetriever(embedder eembedding.Embedder, datasetId int64, reranker *OpenAIChatModel) *HybridRetriever { - return &HybridRetriever{embedder: embedder, datasetId: datasetId, reranker: reranker} +// NewHybridRetriever 构建混合检索器并读取数据集召回数量配置(0=全局默认,-1=尽量多,>0=固定值); +// 各阶段最终取 max(自身值, 最终返回数),保证召回漏斗不断(原始召回/重排候选至少覆盖最终返回数) +func NewHybridRetriever(ctx context.Context, embedder eembedding.Embedder, datasetId int64, reranker *OpenAIChatModel) *HybridRetriever { + r := &HybridRetriever{ + embedder: embedder, + datasetId: datasetId, + reranker: reranker, + vecTopK: consts.VectorTopK, + ftsTopK: consts.FtsTopK, + rerankTopK: consts.RerankTopK, + finalTopK: consts.HybridTopK, + } + if p, err := dao.Dataset.GetRecallParams(ctx, datasetId); err == nil && p != nil { + if p.VecTopK == -1 { + r.vecTopK = consts.MaxRecallRawTopK + } else if p.VecTopK > 0 { + r.vecTopK = p.VecTopK + } + if p.FtsTopK == -1 { + r.ftsTopK = consts.MaxRecallRawTopK + } else if p.FtsTopK > 0 { + r.ftsTopK = p.FtsTopK + } + if p.RerankTopK == -1 { + r.rerankTopK = consts.MaxRerankTopK + } else if p.RerankTopK > 0 { + r.rerankTopK = p.RerankTopK + } + if p.RecallTopK == -1 { + r.finalUnlimited = true + } else if p.RecallTopK > 0 { + r.finalTopK = p.RecallTopK + } + } + if !r.finalUnlimited { + if r.vecTopK < r.finalTopK { + r.vecTopK = r.finalTopK + } + if r.ftsTopK < r.finalTopK { + r.ftsTopK = r.finalTopK + } + if r.rerankTopK < r.finalTopK { + r.rerankTopK = r.finalTopK + } + } + return r } func (r *HybridRetriever) Retrieve(ctx context.Context, query string, opts ...eretriever.Option) ([]*schema.Document, error) { o := eretriever.GetCommonOptions(nil, opts...) - topK := consts.HybridTopK + topK := r.finalTopK if o.TopK != nil && *o.TopK > 0 { topK = *o.TopK } @@ -284,8 +445,8 @@ func (r *HybridRetriever) Retrieve(ctx context.Context, query string, opts ...er items = append(items, scoredChunk{id: id, score: s, sources: srcs[id]}) } sort.Slice(items, func(i, j int) bool { return items[i].score > items[j].score }) - if len(items) > consts.RerankTopK { - items = items[:consts.RerankTopK] + if len(items) > r.rerankTopK { + items = items[:r.rerankTopK] } if r.reranker != nil && len(items) > 0 { before := len(items) @@ -294,26 +455,49 @@ func (r *HybridRetriever) Retrieve(ctx context.Context, query string, opts ...er items[i].score = scores[items[i].id] } sort.Slice(items, func(i, j int) bool { return items[i].score > items[j].score }) - // 门槛作用在重排语义分上:最高分条目必留;其余需同时满足相对比例与绝对下限, - // 防止重排器对泛化条款给出宽松低分(如 3 分)也能进引用 - maxScore := items[0].score - if maxScore > 0 { - floor := math.Max(maxScore*consts.RerankKeepRatio, consts.RerankMinScore) - keep := items[:0] - for i, it := range items { - if i == 0 || it.score >= floor { - keep = append(keep, it) + // 门槛只作用于固定/默认模式:最高分条目必留;其余需同时满足相对比例与绝对下限, + // 防止重排器对泛化条款给出宽松低分(如 3 分)也能进引用。 + // -1 全部召回模式不设分数门槛(用户明确要全部),重排只用于排序 + if !r.finalUnlimited { + maxScore := items[0].score + if maxScore > 0 { + floor := math.Max(maxScore*consts.RerankKeepRatio, consts.RerankMinScore) + keep := items[:0] + for i, it := range items { + if i == 0 || it.score >= floor { + keep = append(keep, it) + } } + items = keep + g.Log().Infof(ctx, "rerank done: %d candidates → %d kept (max %.1f, floor %.1f)", + before, len(items), maxScore, floor) } - items = keep - g.Log().Infof(ctx, "rerank done: %d candidates → %d kept (max %.1f, floor %.1f)", - before, len(items), maxScore, floor) + } else { + g.Log().Infof(ctx, "rerank done: %d candidates, unlimited mode keep all (max %.1f)", + len(items), items[0].score) } } else { g.Log().Warningf(ctx, "rerank failed, fallback to rrf order: %v", err) } } - if len(items) > topK { + if r.finalUnlimited { + // -1 = 不限制条数:按相关性门槛过滤后的条目全返回,仅受总字符预算保护(防上下文爆炸) + total := 0 + keep := items[:0] + for _, it := range items { + chunk, err := dao.Chunk.GetOne(ctx, it.id) + if err != nil || chunk == nil { + continue + } + if total+len([]rune(chunk.Content)) > consts.MaxRecallChars { + g.Log().Infof(ctx, "final unlimited exceeded %d chars, truncated to %d items", consts.MaxRecallChars, len(keep)) + break + } + total += len([]rune(chunk.Content)) + keep = append(keep, it) + } + items = keep + } else if len(items) > topK { items = items[:topK] } @@ -366,7 +550,7 @@ func (r *HybridRetriever) vecRetrieve(ctx context.Context, query string) []retri if len(vecs) == 0 { return hits } - res, err := dao.Chunk.VecSearch(ctx, r.datasetId, domain.VecJsonF64(vecs[0]), consts.VectorTopK) + res, err := dao.Chunk.VecSearch(ctx, r.datasetId, domain.VecJsonF64(vecs[0]), r.vecTopK) if err != nil { g.Log().Warningf(ctx, "vec search failed: %v", err) return hits @@ -379,7 +563,7 @@ func (r *HybridRetriever) vecRetrieve(ctx context.Context, query string) []retri // ftsRetrieve 全文检索段(纯读,供池内并发调用) func (r *HybridRetriever) ftsRetrieve(ctx context.Context, query string) []retrieveHit { - res, err := dao.Chunk.FtsSearch(ctx, r.datasetId, common.TokenizeQuery(query), consts.FtsTopK) + res, err := dao.Chunk.FtsSearch(ctx, r.datasetId, common.TokenizeQuery(query), r.ftsTopK) if err != nil { g.Log().Warningf(ctx, "fts search failed: %v", err) return nil @@ -448,9 +632,34 @@ type chatService struct{} // MaxHistoryRounds 携带进模型的历史对话轮数(每条消息算一条,含用户与助手) const MaxHistoryRounds = 10 -// Ask RAG 问答工作流:混合检索 → 组装提示(含引用编号)→ 对话模型流式生成。 -// history 需已包含最新一条用户问题;onCitations 在检索完成后先于流式输出回调;onDelta 接收增量文本,均可为 nil。 -func (s *chatService) Ask(ctx context.Context, datasetId int64, question string, history []*schema.Message, onCitations func([]domain.Citation), onDelta func(string)) (string, []domain.Citation, error) { +// Ask RAG 问答工作流:数据集配置了 ReAct 轮次时走智能体工具循环(askAgent), +// 否则走单次管线:混合检索 → 组装提示(含引用编号)→ 对话模型流式生成。 +// history 需已包含最新一条用户问题;onCitations 在检索完成后先于流式输出回调; +// onDelta 接收增量文本;onThinking 接收工具轮进度提示,均可为 nil。 +func (s *chatService) Ask(ctx context.Context, datasetId int64, question string, history []*schema.Message, onCitations func([]domain.Citation), onDelta func(string), onThinking func(string)) (string, []domain.Citation, error) { + reactRounds, err := dao.Dataset.GetReactRounds(ctx, datasetId) + if err != nil { + g.Log().Warningf(ctx, "get react rounds failed, fallback to single-pass: %v", err) + } + if reactRounds > 0 { + defaultChatModel, err := dao.ModelConfig.GetDefault(ctx, consts.ModelTypeChat) + if err != nil { + return "", nil, err + } + if defaultChatModel <= 0 { + return "", nil, gerror.New("请先在设置中为对话模型设置默认") + } + model, err := BuildChatModel(ctx, defaultChatModel) + if err != nil { + return "", nil, err + } + msgs := make([]*schema.Message, 0, len(history)+1) + if start := len(history) - MaxHistoryRounds*2; start > 0 { + history = history[start:] + } + msgs = append(msgs, history...) + return s.askAgent(ctx, model, datasetId, reactRounds, question, msgs, onCitations, onDelta, onThinking) + } // 检索(较重,内部再并行 vec/fts)与图增强互不依赖,并行执行;检索放 common.ChatPool,图增强主 goroutine 直接跑 type askOut struct { docs []*schema.Document @@ -481,7 +690,7 @@ func (s *chatService) Ask(ctx context.Context, datasetId int64, question string, return "", nil, out.err } docs := out.docs - citations := buildCitations(docs) + citations := buildCitations(docs, question) if onCitations != nil { onCitations(citations) } @@ -546,11 +755,255 @@ func (s *chatService) retrieve(ctx context.Context, datasetId int64, question st g.Log().Warningf(ctx, "build reranker failed, skip rerank: %v", err) } } - return NewHybridRetriever(emb, datasetId, reranker).Retrieve(ctx, question) + return NewHybridRetriever(ctx, emb, datasetId, reranker).Retrieve(ctx, question) } -// buildCitations 从检索结果生成引用列表(编号从 1 开始,与提示词 [编号] 对应) -func buildCitations(docs []*schema.Document) []domain.Citation { +// askAgent ReAct 循环:模型通过原生 function calling 调用 search 工具检索,最多 rounds 轮; +// 每轮内容增量实时流出,轮次耗尽仍有待执行工具时用聚合上下文强制收尾回答。 +func (s *chatService) askAgent(ctx context.Context, model *OpenAIChatModel, datasetId int64, rounds int, + question string, msgs []*schema.Message, onCitations func([]domain.Citation), onDelta func(string), onThinking func(string)) (string, []domain.Citation, error) { + if err := model.BindTools([]*schema.ToolInfo{searchToolInfo()}); err != nil { + return "", nil, err + } + baseMsgs := append([]*schema.Message(nil), msgs...) + msgs = append(msgs, &schema.Message{Role: schema.System, Content: agentSystemPrompt()}) + + var allDocs []*schema.Document + var allTriples []string + + loopRound := 0 + for { + loopRound++ + if onCitations != nil && len(allDocs) > 0 { + onCitations(aggregateCitations(allDocs, question)) + } + sr, err := model.Stream(ctx, msgs) + if err != nil { + return "", nil, gerror.Wrap(err, "调用对话模型失败") + } + var roundContent strings.Builder + var roundToolCalls []schema.ToolCall + for { + m, err := sr.Recv() + if errors.Is(err, io.EOF) { + break + } + if err != nil { + sr.Close() + return "", nil, gerror.Wrap(err, "流式输出中断") + } + if len(m.ToolCalls) > 0 { + roundToolCalls = append(roundToolCalls, m.ToolCalls...) + continue + } + if m.Content != "" { + roundContent.WriteString(m.Content) + if onDelta != nil { + onDelta(m.Content) + } + } + } + sr.Close() + + if len(roundToolCalls) == 0 { + return roundContent.String(), aggregateCitations(allDocs, question), nil + } + if onThinking != nil { + onThinking("正在检索资料…") + } + msgs = s.execToolCalls(ctx, msgs, roundContent.String(), roundToolCalls, datasetId, &allDocs, &allTriples) + if loopRound >= rounds { + break + } + } + + // 轮次耗尽:执行完最后一批工具调用后,非工具流式收尾(上下文=聚合引用+三元组) + if onCitations != nil { + onCitations(aggregateCitations(allDocs, question)) + } + if onThinking != nil { + onThinking("正在整理答案…") + } + model.BindTools(nil) + finalMsgs := append([]*schema.Message{ + {Role: schema.System, Content: buildSystemPrompt(aggregateCitations(allDocs, question), allTriples)}, + }, baseMsgs...) + sr, err := model.Stream(ctx, finalMsgs) + if err != nil { + return "", nil, gerror.Wrap(err, "调用对话模型失败") + } + defer sr.Close() + var full strings.Builder + for { + m, err := sr.Recv() + if errors.Is(err, io.EOF) { + break + } + if err != nil { + return "", nil, gerror.Wrap(err, "流式输出中断") + } + if m.Content != "" { + full.WriteString(m.Content) + if onDelta != nil { + onDelta(m.Content) + } + } + } + return full.String(), aggregateCitations(allDocs, question), nil +} + +// execToolCalls 执行模型发出的工具调用:解析参数 → 检索 → 追加 assistant/tool 消息并聚合结果。 +// 单个工具失败不中断循环,错误以 JSON 形式作为工具结果返回,模型可自行恢复。 +func (s *chatService) execToolCalls(ctx context.Context, msgs []*schema.Message, roundContent string, + calls []schema.ToolCall, datasetId int64, allDocs *[]*schema.Document, allTriples *[]string) []*schema.Message { + msgs = append(msgs, &schema.Message{Role: schema.Assistant, Content: roundContent, ToolCalls: calls}) + for _, tc := range calls { + var result []byte + query, err := parseSearchQuery(tc.Function.Arguments) + if err != nil { + g.Log().Warningf(ctx, "invalid search tool call: %v", err) + result, _ = json.Marshal(map[string]string{"error": err.Error()}) + } else { + docs, triples, err := s.executeSearchTool(ctx, datasetId, query) + if err != nil { + g.Log().Warningf(ctx, "search tool failed: %v", err) + result, _ = json.Marshal(map[string]string{"error": err.Error()}) + } else { + *allDocs = append(*allDocs, docs...) + *allTriples = append(*allTriples, triples...) + result = buildToolSearchResult(docs, triples) + } + } + msgs = append(msgs, &schema.Message{Role: schema.Tool, ToolCallID: tc.ID, Content: string(result)}) + } + return msgs +} + +// executeSearchTool 执行 search 工具:混合检索与图增强并行(复用 Ask 单次管线的并发模式) +func (s *chatService) executeSearchTool(ctx context.Context, datasetId int64, query string) ([]*schema.Document, []string, error) { + type toolOut struct { + docs []*schema.Document + err error + } + ch := make(chan toolOut, 1) + var wg sync.WaitGroup + wg.Add(1) + if err := common.ChatPool.AddWithRecover(ctx, func(ctx context.Context) { + defer wg.Done() + docs, err := s.retrieve(ctx, datasetId, query) + ch <- toolOut{docs: docs, err: err} + }, func(ctx context.Context, e error) { + defer wg.Done() + ch <- toolOut{err: e} + }); err != nil { + wg.Done() + return nil, nil, err + } + triples, graphErr := KgRelationService.GraphEnhance(ctx, datasetId, query) + wg.Wait() + out := <-ch + if out.err != nil { + return nil, nil, out.err + } + if graphErr != nil { + g.Log().Warningf(ctx, "graph enhance failed: %v", graphErr) + } + return out.docs, triples, nil +} + +// aggregateCitations 多轮工具结果聚合引用:按 chunk_id 去重后重排编号 +func aggregateCitations(docs []*schema.Document, question string) []domain.Citation { + seen := map[int64]bool{} + out := make([]domain.Citation, 0, len(docs)) + for _, c := range buildCitations(docs, question) { + if seen[c.ChunkId] { + continue + } + seen[c.ChunkId] = true + out = append(out, c) + } + for i := range out { + out[i].Index = i + 1 + } + return out +} + +type toolResultItem struct { + ChunkId int64 `json:"chunk_id"` + Content string `json:"content"` + Score float64 `json:"score"` + Sources []string `json:"sources"` +} + +type toolSearchResult struct { + Results []toolResultItem `json:"results"` + GraphTriples []string `json:"graph_triples"` +} + +// buildToolSearchResult 检索结果 → 工具返回 JSON(chunk 内容截断以控制上下文体积) +func buildToolSearchResult(docs []*schema.Document, triples []string) []byte { + res := toolSearchResult{GraphTriples: triples} + for _, d := range docs { + if d == nil { + continue + } + content := d.Content + if rs := []rune(content); len(rs) > consts.ToolResultMaxChars { + content = string(rs[:consts.ToolResultMaxChars]) + } + item := toolResultItem{Content: content} + if id, ok := d.MetaData["chunk_id"].(int64); ok { + item.ChunkId = id + } + if sc, ok := d.MetaData["score"].(float64); ok { + item.Score = sc + } + if srcs, ok := d.MetaData["sources"].([]string); ok { + item.Sources = srcs + } + res.Results = append(res.Results, item) + } + b, _ := json.Marshal(res) + return b +} + +// parseSearchQuery 解析 search 工具参数中的 query +func parseSearchQuery(arguments string) (string, error) { + var p struct { + Query string `json:"query"` + } + if err := json.Unmarshal([]byte(arguments), &p); err != nil { + return "", gerror.New("解析 search 工具参数失败: " + arguments) + } + if strings.TrimSpace(p.Query) == "" { + return "", gerror.New("search 工具缺少 query 参数") + } + return strings.TrimSpace(p.Query), nil +} + +// agentSystemPrompt 智能体模式系统提示:引导模型先检索再作答 +func agentSystemPrompt() string { + return "你是本地知识库智能体。回答用户问题前,应先调用 search 工具检索知识库获取相关资料片段;" + + "若检索结果不足以回答问题,可调整关键词再次检索。每次检索后依据结果继续思考," + + "最终回答需在引用处标注 [N],N 为该句所依据的相关条款数量(依据 1 条写 [1],依据 2 条写 [2]),每个引用处独立计数可重复;" + + "禁止把 [N] 当作编号序列递增使用,也禁止把条文的具体编号(如「第二十一条」)或资料序号写进方括号,条文编号在正文中用文字描述。若某句没有资料依据,不要标注 [0],直接说明资料不足。" +} + +// searchToolInfo search 工具定义:每次调用执行一轮混合检索 + 知识图谱图增强 +func searchToolInfo() *schema.ToolInfo { + return &schema.ToolInfo{ + Name: "search", + Desc: "检索本地知识库获取与问题相关的资料片段(含知识图谱实体关系三元组)。" + + "回答依赖知识库事实时调用;一次检索不足可调整关键词再次调用。", + ParamsOneOf: schema.NewParamsOneOfByParams(map[string]*schema.ParameterInfo{ + "query": {Type: schema.String, Desc: "检索关键词或问题,尽量精简", Required: true}, + }), + } +} + +// buildCitations 从检索结果生成引用列表(编号从 1 开始,与提示词 [编号] 对应), +// 并基于问题计算每条引用中最相关段落的高亮偏移 +func buildCitations(docs []*schema.Document, question string) []domain.Citation { cits := make([]domain.Citation, 0, len(docs)) for i, d := range docs { c := domain.Citation{Index: i + 1, Content: d.Content} @@ -566,18 +1019,85 @@ func buildCitations(docs []*schema.Document) []domain.Citation { if srcs, ok := d.MetaData["sources"].([]string); ok { c.Sources = srcs } + c.HighlightStart, c.HighlightEnd = computeHighlight(d.Content, question) cits = append(cits, c) } return cits } +var clauseStartRe = regexp.MustCompile(`(?m)^\s*第[一二三四五六七八九十百千]+条`) + +// computeHighlight 定位引用内容中与问题最相关的段落(条文/句段),返回其在 content 中的 UTF-16 偏移。 +// 打分 = Σ 命中的问题词字数(长词权重高);总分低于 3 视为未命中(单个 2 字段落如"可以/规定"等泛词不标)。 +func computeHighlight(content, question string) (int, int) { + words := make([]string, 0, 8) + for _, w := range strings.Fields(common.Tokenize(question)) { + if len([]rune(w)) >= 2 { + words = append(words, w) + } + } + if len(words) == 0 || content == "" { + return 0, 0 + } + + rs := []rune(content) + var bounds [][]int + seg := func(s, e int) { + for e > s && (rs[e-1] == '\n' || rs[e-1] == ' ' || rs[e-1] == '\t') { + e-- + } + if e > s { + bounds = append(bounds, []int{s, e}) + } + } + matches := clauseStartRe.FindAllStringIndex(content, -1) + if len(matches) > 1 { + for i, m := range matches { + end := len(rs) + if i+1 < len(matches) { + end = utf8.RuneCountInString(content[:matches[i+1][0]]) + } + seg(utf8.RuneCountInString(content[:m[0]]), end) + } + } else { + start := 0 + for i, r := range rs { + if r == '\n' || r == '。' || r == ';' { + seg(start, i+1) + start = i + 1 + } + } + seg(start, len(rs)) + } + + bestS, bestE, bestScore := 0, 0, 0 + for _, b := range bounds { + score := 0 + for _, w := range words { + if strings.Contains(string(rs[b[0]:b[1]]), w) { + score += len([]rune(w)) + } + } + if score > bestScore { + bestS, bestE, bestScore = b[0], b[1], score + } + } + if bestScore < 3 { + return 0, 0 + } + return len(utf16.Encode(rs[:bestS])), len(utf16.Encode(rs[:bestE])) +} + // buildSystemPrompt 系统提示词:引用资料编号 + 检索片段 + 知识图谱三元组(M5 图增强) func buildSystemPrompt(citations []domain.Citation, triples []string) string { var sb strings.Builder sb.WriteString("你是一个本地知识库助手。请仅根据以下资料回答用户问题;若资料不足以回答,请明确说明。") - sb.WriteString("回答引用资料时,在对应位置标注 [编号]。\n\n【资料】\n") + sb.WriteString("回答引用资料时,在对应位置标注 [N],N 为该句所依据的相关条款数量:依据 1 条写 [1],依据 2 条写 [2],依此类推。") + sb.WriteString("每个引用处独立计数,多次引用可重复相同数字;禁止把 [N] 当作编号序列递增使用(如分点作答写 [1][2][3] 是错误示范)。") + sb.WriteString("禁止把条文的具体编号(如「第二十一条」)或资料序号写进方括号,条文编号请在正文中用文字描述。") + sb.WriteString("若某句没有资料依据,不要标注 [0],直接说明资料不足。\n\n【资料】\n") for _, c := range citations { - sb.WriteString(fmt.Sprintf("[%d] %s\n", c.Index, c.Content)) + sb.WriteString(c.Content + "\n") } if len(triples) > 0 { sb.WriteString("\n【知识图谱】以下为与问题相关的实体关系,可辅助回答关系类问题:\n") @@ -600,7 +1120,25 @@ func buildOpenAIMessages(input []*schema.Message) []openAIMessage { if role == "" { role = string(schema.User) } - out = append(out, openAIMessage{Role: role, Content: m.Content}) + om := openAIMessage{Role: role, Content: m.Content} + switch m.Role { + case schema.Tool: + om.Role = "tool" + om.ToolCallID = m.ToolCallID + case schema.Assistant: + if len(m.ToolCalls) > 0 { + calls := make([]openAIToolCall, 0, len(m.ToolCalls)) + for _, tc := range m.ToolCalls { + calls = append(calls, openAIToolCall{ + ID: tc.ID, + Type: tc.Type, + Function: openAIToolFunction{Name: tc.Function.Name, Arguments: tc.Function.Arguments}, + }) + } + om.ToolCalls = calls + } + } + out = append(out, om) } return out } diff --git a/kb/service/dataset_service.go b/kb/service/dataset_service.go index d4a56e2..a4fb8c1 100644 --- a/kb/service/dataset_service.go +++ b/kb/service/dataset_service.go @@ -26,6 +26,14 @@ func (s *datasetService) Save(ctx context.Context, m *entity.Dataset) (int64, er if m.EmbeddingCfgId == 0 { return 0, gerror.New("数据集必须绑定向量模型,请先选择向量模型") } + if m.ReactRounds < 0 || m.ReactRounds > consts.MaxReactRounds { + return 0, gerror.Newf("智能体轮次需在 0~%d 之间(0=关闭)", consts.MaxReactRounds) + } + for _, v := range []int{m.VecTopK, m.FtsTopK, m.RerankTopK, m.RecallTopK} { + if v < -1 || v > consts.MaxRecallTopK { + return 0, gerror.Newf("召回数量需在 -1~%d 之间(-1=尽量多,0=全局默认)", consts.MaxRecallTopK) + } + } if m.Id > 0 { old, err := dao.Dataset.GetOne(ctx, m.Id) if err != nil { diff --git a/kb/service/message_service.go b/kb/service/message_service.go index e3378e0..a26568d 100644 --- a/kb/service/message_service.go +++ b/kb/service/message_service.go @@ -21,8 +21,8 @@ func (s *messageService) List(ctx context.Context, conversationId int64) ([]*ent } // Chat RAG 问答:会话解析 → 用户消息落库 → 工作流流式生成 → 助手消息+引用落库。 -// onCitations 在检索完成后回调(先于流式输出);onDelta 接收模型增量文本。 -func (s *messageService) Chat(ctx context.Context, conversationId, datasetId int64, question string, onCitations func([]domain.Citation, int64), onDelta func(string)) (string, []domain.Citation, int64, error) { +// onCitations 在检索完成后回调(先于流式输出);onDelta 接收模型增量文本;onThinking 接收工具轮进度提示。 +func (s *messageService) Chat(ctx context.Context, conversationId, datasetId int64, question string, onCitations func([]domain.Citation, int64), onDelta func(string), onThinking func(string)) (string, []domain.Citation, int64, error) { if conversationId <= 0 { title := question if r := []rune(title); len(r) > 20 { @@ -70,7 +70,7 @@ func (s *messageService) Chat(ctx context.Context, conversationId, datasetId int if onCitations != nil { onCitations(c, conversationId) } - }, onDelta) + }, onDelta, onThinking) if err != nil { return "", nil, 0, err } diff --git a/kb/service/system_config_service.go b/kb/service/system_config_service.go index 053a302..3c0b35c 100644 --- a/kb/service/system_config_service.go +++ b/kb/service/system_config_service.go @@ -28,22 +28,29 @@ func (s *systemConfigService) Login(ctx context.Context, token string) (string, return common.SignToken("owner", common.AccessTokenFingerprint(), common.TokenExpireSeconds) } -// GetSettings 读取全局分块默认值(未设置时用内置默认值) -func (s *systemConfigService) GetSettings(ctx context.Context) (chunkSize, chunkOverlap int, err error) { +// GetSettings 读取全局分块默认值与智能体轮次默认值(未设置时用内置默认值) +func (s *systemConfigService) GetSettings(ctx context.Context) (chunkSize, chunkOverlap, reactRounds int, err error) { return dao.AppConfig.GetInt(ctx, consts.SettingsKeyChunkSize, consts.DefaultChunkSize), - dao.AppConfig.GetInt(ctx, consts.SettingsKeyChunkOverlap, consts.DefaultChunkOverlap), nil + dao.AppConfig.GetInt(ctx, consts.SettingsKeyChunkOverlap, consts.DefaultChunkOverlap), + dao.AppConfig.GetInt(ctx, consts.SettingsKeyReactRounds, 0), nil } -// SaveSettings 保存全局分块默认值 -func (s *systemConfigService) SaveSettings(ctx context.Context, chunkSize, chunkOverlap int) error { +// SaveSettings 保存全局分块默认值与智能体轮次默认值 +func (s *systemConfigService) SaveSettings(ctx context.Context, chunkSize, chunkOverlap, reactRounds int) error { if chunkSize < 50 || chunkSize > 5000 { return gerror.New("分块大小需在 50~5000 之间") } if chunkOverlap < 0 || chunkOverlap > 500 { return gerror.New("重叠字数需在 0~500 之间") } + if reactRounds < 0 || reactRounds > consts.MaxReactRounds { + return gerror.Newf("智能体轮次需在 0~%d 之间(0=关闭)", consts.MaxReactRounds) + } if err := dao.AppConfig.SetInt(ctx, consts.SettingsKeyChunkSize, chunkSize); err != nil { return err } - return dao.AppConfig.SetInt(ctx, consts.SettingsKeyChunkOverlap, chunkOverlap) + if err := dao.AppConfig.SetInt(ctx, consts.SettingsKeyChunkOverlap, chunkOverlap); err != nil { + return err + } + return dao.AppConfig.SetInt(ctx, consts.SettingsKeyReactRounds, reactRounds) } diff --git a/ui-src/src/api/chat.js b/ui-src/src/api/chat.js index 5e0f33c..c0509ad 100644 --- a/ui-src/src/api/chat.js +++ b/ui-src/src/api/chat.js @@ -61,6 +61,7 @@ export async function streamChat(payload, handlers) { } if (obj.citations !== undefined) handlers.onCitations(obj) else if (obj.content !== undefined) handlers.onDelta(obj.content) + else if (obj.type === 'thinking' && handlers.onThinking) handlers.onThinking(obj.message || '') else if (obj.status === 'ok') handlers.onDone() else if (obj.message) handlers.onError(new Error(obj.message)) } diff --git a/ui-src/src/views/Chat.vue b/ui-src/src/views/Chat.vue index bb4d8b9..d82a583 100644 --- a/ui-src/src/views/Chat.vue +++ b/ui-src/src/views/Chat.vue @@ -25,6 +25,7 @@