Files
rag-local/kb/service/chat_service.go
T
2026-08-07 16:07:20 +08:00

1192 lines
38 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package service
import (
"bufio"
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"math"
"net/http"
"regexp"
"sort"
"strconv"
"strings"
"sync"
"time"
"unicode/utf16"
"unicode/utf8"
"rag-local/common"
"rag-local/kb/consts"
"rag-local/kb/dao"
"rag-local/kb/model/domain"
"rag-local/kb/model/entity"
eembedding "github.com/cloudwego/eino/components/embedding"
emodel "github.com/cloudwego/eino/components/model"
eretriever "github.com/cloudwego/eino/components/retriever"
"github.com/cloudwego/eino/schema"
"github.com/gogf/gf/v2/errors/gerror"
"github.com/gogf/gf/v2/frame/g"
)
var httpClient = &http.Client{Timeout: 2 * time.Minute}
// ---------- OpenAI 兼容 HTTP 组件 ----------
type openAIMessage struct {
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.ToolCallIndex 仅流式累积时使用)
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 {
Choices []struct {
Message openAIMessage `json:"message"`
FinishReason string `json:"finish_reason"`
} `json:"choices"`
Usage *schema.TokenUsage `json:"usage"`
}
type openAIStreamChunk struct {
Choices []struct {
Delta openAIMessage `json:"delta"`
FinishReason string `json:"finish_reason"`
} `json:"choices"`
}
// OpenAIChatModel 基于 OpenAI 兼容 /chat/completions 接口的对话模型,实现 eino model.ChatModel
type OpenAIChatModel struct {
cfg *entity.ModelConfig
tools []*schema.ToolInfo
}
func NewOpenAIChatModel(cfg *entity.ModelConfig) *OpenAIChatModel {
return &OpenAIChatModel{cfg: cfg}
}
func (m *OpenAIChatModel) Generate(ctx context.Context, input []*schema.Message, opts ...emodel.Option) (*schema.Message, error) {
payload := map[string]any{
"model": m.cfg.ModelName,
"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
}
var resp openAIChatResponse
if err := json.Unmarshal(body, &resp); err != nil {
return nil, gerror.Wrap(err, "解析模型响应失败")
}
if len(resp.Choices) == 0 {
return nil, gerror.New("模型返回空响应")
}
choice := resp.Choices[0]
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}
}
return msg, nil
}
func (m *OpenAIChatModel) Stream(ctx context.Context, input []*schema.Message, opts ...emodel.Option) (*schema.StreamReader[*schema.Message], error) {
payload := map[string]any{
"model": m.cfg.ModelName,
"messages": buildOpenAIMessages(input),
"stream": true,
}
m.withTools(payload)
reader, writer := schema.Pipe[*schema.Message](16)
go func() {
defer writer.Close()
body, err := postOpenAIStream(ctx, m.cfg, m.endpoint("/chat/completions"), payload)
if err != nil {
writer.Send(nil, err)
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')
text := strings.TrimSpace(string(line))
if strings.HasPrefix(text, "data:") {
data := strings.TrimSpace(strings.TrimPrefix(text, "data:"))
if data == "[DONE]" {
break
}
var chunk openAIStreamChunk
if json.Unmarshal([]byte(data), &chunk) == nil && len(chunk.Choices) > 0 {
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(去除 IndexType 缺省 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
}
// OpenAIEmbedder 基于 OpenAI 兼容 /embeddings 接口的向量模型,实现 eino embedding.Embedder
type OpenAIEmbedder struct {
cfg *entity.ModelConfig
}
func NewOpenAIEmbedder(cfg *entity.ModelConfig) *OpenAIEmbedder {
return &OpenAIEmbedder{cfg: cfg}
}
// Dim 配置的向量维度(vec0 表建表维度需一致)
func (e *OpenAIEmbedder) Dim() int {
if e.cfg.Dimension > 0 {
return e.cfg.Dimension
}
return consts.DefaultEmbeddingDim
}
// EmbedStrings 内部分批请求(dashscope 单次最多 20 条),返回与 texts 同序的向量
func (e *OpenAIEmbedder) EmbedStrings(ctx context.Context, texts []string, opts ...eembedding.Option) ([][]float64, error) {
out := make([][]float64, len(texts))
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
}
// BuildChatModel 按配置 id 构建对话模型
func BuildChatModel(ctx context.Context, cfgId int64) (*OpenAIChatModel, error) {
cfg, err := dao.ModelConfig.GetOne(ctx, cfgId)
if err != nil {
return nil, err
}
if cfg == nil {
return nil, gerror.New("模型配置不存在")
}
if cfg.ModelType != consts.ModelTypeChat {
return nil, gerror.New("该模型配置不是对话模型(model_type=chat")
}
if cfg.EndpointUrl == "" || cfg.ModelName == "" {
return nil, gerror.New("模型配置缺少 endpoint_url 或 model_name")
}
return NewOpenAIChatModel(cfg), nil
}
// BuildEmbedder 按配置 id 构建向量模型
func BuildEmbedder(ctx context.Context, cfgId int64) (*OpenAIEmbedder, error) {
cfg, err := dao.ModelConfig.GetOne(ctx, cfgId)
if err != nil {
return nil, err
}
if cfg == nil {
return nil, gerror.New("模型配置不存在")
}
if cfg.ModelType != consts.ModelTypeEmbedding {
return nil, gerror.New("该模型配置不是向量模型(model_type=embedding")
}
if cfg.EndpointUrl == "" || cfg.ModelName == "" {
return nil, gerror.New("模型配置缺少 endpoint_url 或 model_name")
}
return NewOpenAIEmbedder(cfg), nil
}
// HybridRetriever 混合检索器:向量 KNN + FTS5 BM25RRF 融合后可选 LLM 重排,实现 eino retriever.Retriever
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
}
// 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 := r.finalTopK
if o.TopK != nil && *o.TopK > 0 {
topK = *o.TopK
}
scores := make(map[int64]float64)
srcs := make(map[int64][]string)
// 向量段与全文段互不依赖,提交 common.ChatRetrievePool 并行执行,RRF 合并回主 goroutine 串行
ch := make(chan []retrieveHit, 2)
var wg sync.WaitGroup
if r.embedder != nil {
wg.Add(1)
if err := common.ChatRetrievePool.AddWithRecover(ctx, func(ctx context.Context) {
defer wg.Done()
ch <- r.vecRetrieve(ctx, query)
}, func(ctx context.Context, e error) {
defer wg.Done()
ch <- nil
g.Log().Warningf(ctx, "vec retrieve failed: %v", e)
}); err != nil {
wg.Done()
g.Log().Warningf(ctx, "submit vec retrieve failed: %v", err)
}
}
wg.Add(1)
if err := common.ChatRetrievePool.AddWithRecover(ctx, func(ctx context.Context) {
defer wg.Done()
ch <- r.ftsRetrieve(ctx, query)
}, func(ctx context.Context, e error) {
defer wg.Done()
ch <- nil
g.Log().Warningf(ctx, "fts retrieve failed: %v", e)
}); err != nil {
wg.Done()
g.Log().Warningf(ctx, "submit fts retrieve failed: %v", err)
}
go func() { wg.Wait(); close(ch) }()
for hits := range ch {
for _, h := range hits {
scores[h.chunkId] += 1 / (float64(consts.RrfK) + float64(h.rank) + 1)
srcs[h.chunkId] = append(srcs[h.chunkId], h.source)
}
}
items := make([]scoredChunk, 0, len(scores))
for id, s := range scores {
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) > r.rerankTopK {
items = items[:r.rerankTopK]
}
if r.reranker != nil && len(items) > 0 {
before := len(items)
if scores, err := r.rerankByLLM(ctx, query, items); err == nil {
for i := range items {
items[i].score = scores[items[i].id]
}
sort.Slice(items, func(i, j int) bool { return items[i].score > items[j].score })
// 门槛只作用于固定/默认模式:最高分条目必留;其余需同时满足相对比例与绝对下限,
// 防止重排器对泛化条款给出宽松低分(如 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)
}
} 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 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]
}
docs := make([]*schema.Document, 0, len(items))
for _, it := range items {
chunk, err := dao.Chunk.GetOne(ctx, it.id)
if err != nil || chunk == nil {
continue
}
docs = append(docs, &schema.Document{
ID: strconv.FormatInt(chunk.Id, 10),
Content: chunk.Content,
MetaData: map[string]any{
"chunk_id": chunk.Id,
"document_id": chunk.DocumentId,
"seq": chunk.Seq,
"score": it.score,
"sources": it.sources,
},
})
}
return docs, nil
}
// scoredChunk RRF 融合后的候选条目(重排后 score 字段替换为语义分)
type scoredChunk struct {
id int64
score float64
sources []string
}
// retrieveHit 单路检索命中(rank 用于 RRF 融合)
type retrieveHit struct {
chunkId int64
rank int
source string
}
// vecRetrieve 向量段检索(纯读,供池内并发调用)
func (r *HybridRetriever) vecRetrieve(ctx context.Context, query string) []retrieveHit {
var hits []retrieveHit
if r.embedder == nil {
return hits
}
vecs, err := r.embedder.EmbedStrings(ctx, []string{query})
if err != nil {
g.Log().Warningf(ctx, "query embed failed: %v", err)
return hits
}
if len(vecs) == 0 {
return hits
}
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
}
for i, h := range res {
hits = append(hits, retrieveHit{chunkId: h.ChunkId, rank: i, source: "vector"})
}
return hits
}
// ftsRetrieve 全文检索段(纯读,供池内并发调用)
func (r *HybridRetriever) ftsRetrieve(ctx context.Context, query string) []retrieveHit {
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
}
hits := make([]retrieveHit, 0, len(res))
for i, h := range res {
hits = append(hits, retrieveHit{chunkId: h.ChunkId, rank: i, source: "fts"})
}
return hits
}
// rerankByLLM 用默认对话模型对候选分块打分(0-10,JSON 输出),返回 chunk_id → 相关性分;
// 任何失败(调用/解析/空结果)返回错误,由调用方回退 RRF 排序
func (r *HybridRetriever) rerankByLLM(ctx context.Context, query string, items []scoredChunk) (map[int64]float64, error) {
var sb strings.Builder
sb.WriteString("你是检索重排器。请评估每个候选段落与用户问题的相关性,为每个候选输出 0-10 的相关性分数(10=高度相关,0=完全不相关)。\n\n用户问题:\n")
sb.WriteString(query)
sb.WriteString("\n\n候选段落:\n")
for i, it := range items {
chunk, err := dao.Chunk.GetOne(ctx, it.id)
if err != nil || chunk == nil {
continue
}
content := chunk.Content
if rs := []rune(content); len(rs) > consts.RerankMaxChars {
content = string(rs[:consts.RerankMaxChars])
}
sb.WriteString(fmt.Sprintf("[%d] %s\n", i+1, content))
}
sb.WriteString("\n只输出 JSON,不要其他内容:{\"scores\":{\"1\":8,\"2\":3}}")
msg, err := r.reranker.Generate(ctx, []*schema.Message{{Role: schema.User, Content: sb.String()}})
if err != nil {
return nil, err
}
// 模型可能包裹 ```json 代码块,截取首尾花括号之间的 JSON 对象
content := msg.Content
if i, j := strings.Index(content, "{"), strings.LastIndex(content, "}"); i >= 0 && j > i {
content = content[i : j+1]
}
var resp struct {
Scores map[string]float64 `json:"scores"`
}
if err := json.Unmarshal([]byte(content), &resp); err != nil {
return nil, gerror.Wrap(err, "解析重排结果失败: "+msg.Content)
}
scores := make(map[int64]float64, len(resp.Scores))
for k, v := range resp.Scores {
idx, err := strconv.ParseInt(k, 10, 64)
if err != nil || idx < 1 || idx > int64(len(items)) {
continue
}
scores[items[idx-1].id] = v
}
if len(scores) == 0 {
return nil, gerror.New("重排结果为空")
}
return scores, nil
}
// ---------- RAG 问答工作流 ----------
var ChatService = &chatService{}
type chatService struct{}
// MaxHistoryRounds 携带进模型的历史对话轮数(每条消息算一条,含用户与助手)
const MaxHistoryRounds = 10
// 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
err error
}
ch := make(chan askOut, 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, question)
ch <- askOut{docs: docs, err: err}
}, func(ctx context.Context, e error) {
defer wg.Done()
ch <- askOut{err: e}
}); err != nil {
wg.Done()
return "", nil, err
}
triples, err := KgRelationService.GraphEnhance(ctx, datasetId, question)
if err != nil {
g.Log().Warningf(ctx, "graph enhance failed: %v", err)
}
wg.Wait()
out := <-ch
if out.err != nil {
return "", nil, out.err
}
docs := out.docs
citations := buildCitations(docs, question)
if onCitations != nil {
onCitations(citations)
}
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)
msgs = append(msgs, &schema.Message{Role: schema.System, Content: buildSystemPrompt(citations, triples)})
if start := len(history) - MaxHistoryRounds*2; start > 0 {
history = history[start:]
}
msgs = append(msgs, history...)
sr, err := model.Stream(ctx, msgs)
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, "流式输出中断")
}
full.WriteString(m.Content)
if onDelta != nil {
onDelta(m.Content)
}
}
return full.String(), citations, nil
}
// retrieve 构建数据集绑定的混合检索器并执行检索(无 embedding 配置时仅全文)
func (s *chatService) retrieve(ctx context.Context, datasetId int64, question string) ([]*schema.Document, error) {
var emb eembedding.Embedder
if cfgId, err := dao.Dataset.GetEmbeddingCfgId(ctx, datasetId); err == nil && cfgId > 0 {
if em, err := BuildEmbedder(ctx, cfgId); err == nil {
emb = em
} else {
g.Log().Warningf(ctx, "build embedder failed, retrieve fts only: %v", err)
}
}
// LLM 重排器:默认对话模型;构建失败仅跳过重排,不影响检索主流程
var reranker *OpenAIChatModel
if defaultChatModel, err := dao.ModelConfig.GetDefault(ctx, consts.ModelTypeChat); err == nil && defaultChatModel > 0 {
if m, err := BuildChatModel(ctx, defaultChatModel); err == nil {
reranker = m
} else {
g.Log().Warningf(ctx, "build reranker failed, skip rerank: %v", err)
}
}
return NewHybridRetriever(ctx, emb, datasetId, reranker).Retrieve(ctx, question)
}
// 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}
if id, ok := d.MetaData["chunk_id"].(int64); ok {
c.ChunkId = id
}
if id, ok := d.MetaData["document_id"].(int64); ok {
c.DocumentId = id
}
if s, ok := d.MetaData["score"].(float64); ok {
c.Score = s
}
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 为该句所依据的相关条款数量:依据 1 条写 [1],依据 2 条写 [2],依此类推。")
sb.WriteString("每个引用处独立计数,多次引用可重复相同数字;禁止把 [N] 当作编号序列递增使用(如分点作答写 [1][2][3] 是错误示范)。")
sb.WriteString("禁止把条文的具体编号(如「第二十一条」)或资料序号写进方括号,条文编号请在正文中用文字描述。")
sb.WriteString("若某句没有资料依据,不要标注 [0],直接说明资料不足。\n\n【资料】\n")
for _, c := range citations {
sb.WriteString(c.Content + "\n")
}
if len(triples) > 0 {
sb.WriteString("\n【知识图谱】以下为与问题相关的实体关系,可辅助回答关系类问题:\n")
for _, t := range triples {
sb.WriteString(t + "\n")
}
}
return sb.String()
}
// ---------- HTTP 辅助 ----------
func buildOpenAIMessages(input []*schema.Message) []openAIMessage {
out := make([]openAIMessage, 0, len(input))
for _, m := range input {
if m == nil {
continue
}
role := string(m.Role)
if role == "" {
role = string(schema.User)
}
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
}
func postOpenAI(ctx context.Context, cfg *entity.ModelConfig, url string, payload any) ([]byte, error) {
body, err := doOpenAIRequest(ctx, cfg, url, payload)
if err != nil {
return nil, err
}
defer body.Close()
resp, err := io.ReadAll(body)
if err != nil {
return nil, err
}
return resp, nil
}
func postOpenAIStream(ctx context.Context, cfg *entity.ModelConfig, url string, payload any) (io.ReadCloser, error) {
return doOpenAIRequest(ctx, cfg, url, payload)
}
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
}
req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(buf))
if err != nil {
return nil, err
}
req.Header.Set("Content-Type", "application/json")
if cfg.ApiKey != "" {
req.Header.Set("Authorization", "Bearer "+cfg.ApiKey)
}
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()
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
}