1277 lines
42 KiB
Go
1277 lines
42 KiB
Go
package service
|
||
|
||
import (
|
||
"bufio"
|
||
"bytes"
|
||
"context"
|
||
"encoding/json"
|
||
"errors"
|
||
"fmt"
|
||
"io"
|
||
"math"
|
||
"net"
|
||
"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"
|
||
)
|
||
|
||
// ---------- 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.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 {
|
||
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
|
||
extra map[string]any // 附加请求字段(如 gemma-4 关闭思维链)
|
||
}
|
||
|
||
func NewOpenAIChatModel(cfg *entity.ModelConfig) *OpenAIChatModel {
|
||
return &OpenAIChatModel{cfg: cfg, extra: map[string]any{}}
|
||
}
|
||
|
||
// DisableThinking 关闭模型思维链:批量任务(标注/图谱抽取)保持快速稳定,避免长思考拖慢甚至超时
|
||
// 经 mlx server 的 chat_template_kwargs 透传模板参数(Qwen3.5 模板键为 enable_thinking)
|
||
func (m *OpenAIChatModel) DisableThinking() *OpenAIChatModel {
|
||
m.extra["chat_template_kwargs"] = map[string]any{"enable_thinking": false}
|
||
return m
|
||
}
|
||
|
||
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,
|
||
}
|
||
for k, v := range m.extra {
|
||
payload[k] = v
|
||
}
|
||
if commonOpts := emodel.GetCommonOptions(nil, opts...); commonOpts.MaxTokens != nil {
|
||
payload["max_tokens"] = *commonOpts.MaxTokens
|
||
}
|
||
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(去除 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
|
||
}
|
||
|
||
// 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 BM25,RRF 融合后可选 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=完全不相关)。")
|
||
sb.WriteString("严格规则:仅当段落直接回答问题的具体情境(如涉及偷窃、归还、退赃等具体事实)时才给高分(7-10);")
|
||
sb.WriteString("段落只是泛化提及相关概念(如犯罪、处罚、共同犯罪等一般规定)而未涉及问题具体情境时,必须给低分(0-3)。\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}}")
|
||
// 重排打分任务关闭思考链并限制输出上限:Qwen3.5 思考链会输出千级 token 的"Thinking Process"
|
||
// 长文(~5 tok/s 下可达数分钟),无 max_tokens 兜底会拖到应用超时;且思考文本混入 content 破坏 JSON 解析
|
||
msg, err := r.reranker.Generate(ctx, []*schema.Message{{Role: schema.User, Content: sb.String()}},
|
||
emodel.WithMaxTokens(consts.RerankMaxTokens))
|
||
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
|
||
}
|
||
model.DisableThinking() // 工具轮次优先稳定快速,思考文本会混入 content 干扰解析
|
||
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
|
||
}
|
||
model.DisableThinking() // Qwen3.5 思考文本会直接混入 content,关闭后输出干净答案
|
||
|
||
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.DisableThinking() // 重排打分无需思考链,思考文本会拖慢任务并破坏 JSON 解析
|
||
} 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
|
||
|
||
// 首轮预检索:不依赖模型自觉调用 search——9B 模型思考链关闭后偶发跳过工具调用直接凭自身知识作答,
|
||
// 导致检索不执行、引用列表为空、回答 [N] 为模型自编。强制注入首轮结果,模型仍可后续补充检索。
|
||
if docs, triples, err := s.executeSearchTool(ctx, datasetId, question); err != nil {
|
||
g.Log().Warningf(ctx, "pre-retrieve failed, fallback to agent search tool: %v", err)
|
||
} else if len(docs) > 0 {
|
||
allDocs = append(allDocs, docs...)
|
||
allTriples = append(allTriples, triples...)
|
||
var sb strings.Builder
|
||
sb.WriteString("已为你检索到以下资料,回答时请优先基于这些资料,并在引用处标注 [N](N 为该句依据的资料条数);若资料不足,可继续调用 search 工具补充检索:\n\n")
|
||
for i, c := range buildCitations(docs, question) {
|
||
sb.WriteString(fmt.Sprintf("[%d] %s\n", i+1, c.Content))
|
||
}
|
||
if len(triples) > 0 {
|
||
sb.WriteString("\n【知识图谱】以下为与问题相关的实体关系,可辅助回答关系类问题:\n")
|
||
for _, t := range triples {
|
||
sb.WriteString(t + "\n")
|
||
}
|
||
}
|
||
msgs = append(msgs, &schema.Message{Role: schema.User, Content: sb.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)
|
||
}
|
||
sort.Slice(out, func(i, j int) bool { return out[i].Score > out[j].Score })
|
||
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 {
|
||
content := c.Content
|
||
if rs := []rune(content); len(rs) > consts.CitationMaxChars {
|
||
content = string(rs[:consts.CitationMaxChars])
|
||
}
|
||
sb.WriteString(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
|
||
}
|
||
// 本地模型生成慢且多请求排队时响应可超分钟级,超时取 config.yml chat.timeout(秒)
|
||
timeout := g.Cfg().MustGet(ctx, "chat.timeout", consts.DefaultLlmHttpTimeout).Int()
|
||
maxRetries := g.Cfg().MustGet(ctx, "chat.max_retries", 0).Int()
|
||
client := &http.Client{Timeout: time.Duration(timeout) * time.Second}
|
||
var lastErr error
|
||
for attempt := 0; attempt <= maxRetries; attempt++ {
|
||
if attempt > 0 {
|
||
g.Log().Warningf(ctx, "model call retry %d/%d: model=%s url=%s err=%v", attempt, maxRetries, cfg.ModelName, url, lastErr)
|
||
select {
|
||
case <-ctx.Done():
|
||
return nil, ctx.Err()
|
||
case <-time.After(time.Duration(1<<(attempt-1)) * time.Second):
|
||
}
|
||
}
|
||
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 := client.Do(req)
|
||
if err != nil {
|
||
lastErr = err
|
||
if attempt < maxRetries && isRetryableLLMErr(err) {
|
||
continue
|
||
}
|
||
g.Log().Errorf(ctx, "model call failed: model=%s url=%s err=%v", cfg.ModelName, url, err)
|
||
return nil, err
|
||
}
|
||
if resp.StatusCode >= 500 && attempt < maxRetries {
|
||
msg, _ := io.ReadAll(resp.Body)
|
||
resp.Body.Close()
|
||
lastErr = gerror.Newf("模型接口 %s 返回 %d: %s", url, resp.StatusCode, strings.TrimSpace(string(msg)))
|
||
continue
|
||
}
|
||
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
|
||
}
|
||
return nil, lastErr
|
||
}
|
||
|
||
// isRetryableLLMErr 是否值得重试:超时/取消类不重试(源于排队或慢响应,重试只会重新排队);
|
||
// 连接类瞬时错误(拒绝/重置/EOF 等)与 5xx 服务端错误可重试
|
||
func isRetryableLLMErr(err error) bool {
|
||
if errors.Is(err, context.DeadlineExceeded) || errors.Is(err, context.Canceled) {
|
||
return false
|
||
}
|
||
var ne net.Error
|
||
if errors.As(err, &ne) && ne.Timeout() {
|
||
return false
|
||
}
|
||
return true
|
||
}
|