Files
pj0235-eai_agentplatform/eai_agentplatform/backend-go/internal/ai/retrieve.go
T
eaiadminandClaude Code c1af86c934 feat(asr): 本地语音转写接入为一级路由 + 并行工作流合并提交
按用户指示做**一包提交**,不按工作流拆分。本提交刻意混合了多条并行线:

  · 本地 ASR 接管:audio 成为与 chat/embed/image/video 同等的路由类别
    (IsLocalRoute 单一判据、audio 健康探测、default_audio_route、
    auto 占位、GET /api/ai/routes/audio、回退云端时界面明示「音频已出网」)
  · LLM 调用层:ctx 贯穿、ToolCall/ToolSchema、EmptyCompletionError /
    TransientUpstreamError(按错误类型而非文案判重试)
  · 编排 Agent:general_assistant orchestrate/persistence/spec_driver
  · 联网搜索:internal/search(playwright)
  · 网盘:backend + 前端
  · 前端 UI:导航/路由/工作台若干页
  · 交付文档:DELIVERY.md / AR04 / 部署文档的「无 Python」表述据实改写,
    新增 eai_agentplatform-asr.service、asr.env、clonezilla-cleanup 清 ~/asr-poc

不分拆的原因:dev 早期,粒度不该打断工作节奏。且实测过——这些改动
**在编译上是同一个单元**(llm.go 的 ctx 签名变更牵动 12 个调用点,
chat_message.go 的 ctx 改动又与编排重写同处一个 hunk),拆出来的中间态编不过。
详见 TOP_CODING_RULES.md G14.5 与 bugs_and_errors.md E09。

Co-Authored-By: Claude Code <noreply@anthropic.com>
2026-09-26 22:21:39 +08:00

165 lines
3.5 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 ai
import (
"context"
"math"
"sort"
"strings"
"eai_agentplatform/backend/internal/config"
"eai_agentplatform/backend/internal/model"
"eai_agentplatform/backend/internal/store"
)
// Retrieve 混合检索:向量 brute-force 余弦 + 关键词兜底,去重合并取 topK
func Retrieve(cfg *config.Config, query string, topK int) []string {
if topK <= 0 {
topK = 5
}
var vectorResults []string
if aiRoute, err := config.GetRoute("embed_gen"); err == nil {
if v, err := vectorRetrieve(aiRoute, query, topK); err == nil {
vectorResults = v
}
}
keywordResults := keywordRetrieve(query, topK*2)
seen := map[string]bool{}
out := make([]string, 0, topK)
for _, s := range append(vectorResults, keywordResults...) {
s = strings.TrimSpace(s)
if s == "" || seen[s] {
continue
}
seen[s] = true
out = append(out, s)
if len(out) >= topK {
break
}
}
return out
}
// vectorRetrieve 向量检索:query + 所有 chunk 一次批量 embedding,brute-force 余弦 topK
func vectorRetrieve(aiRoute *config.RouteConfig, query string, topK int) ([]string, error) {
var chunks []model.KnowledgeChunk
store.DB.Order("id ASC").Find(&chunks)
if len(chunks) == 0 {
return nil, nil
}
client := NewClient(aiRoute)
inputs := make([]string, 0, len(chunks)+1)
inputs = append(inputs, query)
for _, ch := range chunks {
inputs = append(inputs, ch.Content)
}
vecs, err := client.Embed(context.Background(), inputs)
if err != nil || len(vecs) != len(inputs) {
if err != nil {
return nil, err
}
return nil, nil
}
qv := vecs[0]
type scored struct {
idx int
sim float64
}
ss := make([]scored, 0, len(chunks))
for i := 1; i < len(vecs); i++ {
ss = append(ss, scored{i - 1, cosine(qv, vecs[i])})
}
sort.Slice(ss, func(a, b int) bool { return ss[a].sim > ss[b].sim })
seen := map[string]bool{}
out := make([]string, 0, topK)
for _, s := range ss {
content := chunks[s.idx].Content
if seen[content] {
continue
}
seen[content] = true
out = append(out, content)
if len(out) >= topK {
break
}
}
return out, nil
}
// keywordRetrieve 关键词兜底:term 命中数排序
func keywordRetrieve(query string, topK int) []string {
terms := splitTerms(query)
var chunks []model.KnowledgeChunk
store.DB.Order("id ASC").Find(&chunks)
type scored struct {
content string
score int
}
var ss []scored
for _, ch := range chunks {
s := 0
for _, t := range terms {
if strings.Contains(ch.Content, t) {
s++
}
}
if s > 0 {
ss = append(ss, scored{ch.Content, s})
}
}
sort.Slice(ss, func(a, b int) bool { return ss[a].score > ss[b].score })
seen := map[string]bool{}
out := make([]string, 0, topK)
for _, s := range ss {
if seen[s.content] {
continue
}
seen[s.content] = true
out = append(out, s.content)
if len(out) >= topK {
break
}
}
return out
}
func splitTerms(q string) []string {
f := func(r rune) bool {
switch r {
case ' ', ',', '。', '?', '!', '、', ',', '.', '?', '!', ':', ':', ';', ';':
return true
}
return false
}
terms := strings.FieldsFunc(q, f)
var out []string
for _, t := range terms {
if len([]rune(t)) >= 2 {
out = append(out, t)
}
}
if len(out) == 0 {
out = []string{q}
}
return out
}
func cosine(a, b []float64) float64 {
if len(a) == 0 || len(a) != len(b) {
return 0
}
var dot, na, nb float64
for i := range a {
dot += a[i] * b[i]
na += a[i] * a[i]
nb += b[i] * b[i]
}
if na == 0 || nb == 0 {
return 0
}
return dot / (math.Sqrt(na) * math.Sqrt(nb))
}