Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 11 additions & 3 deletions agent/compact.go
Original file line number Diff line number Diff line change
Expand Up @@ -250,7 +250,7 @@ func EstimatePromptTokens(workspace, skillCatalog, summary string, history []Cha
// [lastSystemPrompt] + history[:keepStart] + [尾部压缩指令],并带上 lastToolSpecsJSON 还原的
// 工具集 —— 这串前缀正是上次缓存下来的,几乎全命中,只有尾部指令是 miss。
// lastSystemPrompt 为空(无快照)时退回冷路径:compressionPrompt 当 system + 拍平历史。
func RunCompression(lastSystemPrompt, lastToolSpecsJSON string, history []ChatMessage, entry ModelEntry, ctxWin int) (
func RunCompression(lastSystemPrompt, lastToolSpecsJSON string, history []ChatMessage, entry ModelEntry, ctxWin int, focusHint string) (
summary string, cutIdx int, compressedTurns int, err error) {

// 轮数按 isTurnBoundary(user 或 assistant)计:一个 user 消息 + 几十轮工具调用同样是几十轮对话,
Expand Down Expand Up @@ -332,17 +332,25 @@ func RunCompression(lastSystemPrompt, lastToolSpecsJSON string, history []ChatMe
convo := make([]ChatMessage, 0, keepStart+2)
convo = append(convo, ChatMessage{Role: "system", Content: lastSystemPrompt})
convo = append(convo, history[:keepStart]...)
convo = append(convo, ChatMessage{Role: "user", Content: warmCompressInstruction})
instruction := warmCompressInstruction
if focusHint != "" {
instruction = fmt.Sprintf("%s\n\n**压缩侧重点**: 请重点关注与[%s]相关的内容, 保留相关决策和上下文; 与侧重点无关的内容可以更激进地压缩。", instruction, focusHint)
}
convo = append(convo, ChatMessage{Role: "user", Content: instruction})
toolSpecs := UnmarshalToolSpecs(lastToolSpecsJSON)
summary, err = CallWithTools(ctx, entry.APIKey, entry.BaseURL, entry.Model, convo, toolSpecs, summaryMax)
} else {
// 冷路径:无快照,拍平历史走独立 system(必 miss,但正确)。
cp := compressionPrompt
if focusHint != "" {
cp = fmt.Sprintf("%s\n\n**压缩侧重点**: 请重点关注与[%s]相关的内容, 保留相关决策和上下文; 与侧重点无关的内容可以更激进地压缩。", cp, focusHint)
}
var inputBuf strings.Builder
for _, msg := range history[:keepStart] {
inputBuf.WriteString("[" + msg.Role + "]\n" + msg.Content + "\n\n")
}
convo := []ChatMessage{
{Role: "system", Content: compressionPrompt},
{Role: "system", Content: cp},
{Role: "user", Content: inputBuf.String()},
}
summary, err = CallOnce(ctx, entry.APIKey, entry.BaseURL, entry.Model, convo, summaryMax)
Expand Down
4 changes: 2 additions & 2 deletions agent/compact_cooldown_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,7 @@ func TestRunCompression_TooFewTurnsSentinel(t *testing.T) {
{Role: "user", Content: "一"},
{Role: "assistant", Content: "回一"},
} // 2 轮(user + assistant),不多于要保留的 keepRecentTurns
_, _, _, err := RunCompression("", "", hist, ModelEntry{ContextWindow: 100000}, 100000)
_, _, _, err := RunCompression("", "", hist, ModelEntry{ContextWindow: 100000}, 100000, "")
if !errors.Is(err, ErrCompactTooFewTurns) {
t.Fatalf("2 轮应返回 ErrCompactTooFewTurns 哨兵, got %v", err)
}
Expand All @@ -34,7 +34,7 @@ func TestRunCompression_SingleUserLongTurnNotRejected(t *testing.T) {
hist = append(hist, asstCall(id, "Bash", `{"command":"go test"}`), toolMsg(id, "Bash", body))
}
// BaseURL 为空 → 摘要请求在本地就失败;这里只关心它已越过轮数 / 切点守卫。
_, _, _, err := RunCompression("sys", "[]", hist, ModelEntry{ContextWindow: 20000}, 20000)
_, _, _, err := RunCompression("sys", "[]", hist, ModelEntry{ContextWindow: 20000}, 20000, "")
if errors.Is(err, ErrCompactTooFewTurns) {
t.Fatal("单个 user 长任务轮不应再被判成轮数不足")
}
Expand Down
2 changes: 1 addition & 1 deletion agent/compact_live_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -62,7 +62,7 @@ func TestLive_RunCompressionSucceeds(t *testing.T) {
}
t.Logf("历史 ≈ %d tokens", EstimateHistoryTokens(hist))

summary, cutIdx, turns, err := RunCompression("", "", hist, entry, ctxWin)
summary, cutIdx, turns, err := RunCompression("", "", hist, entry, ctxWin, "")
if err != nil {
t.Fatalf("❌ 真实压缩失败(正常路径不该失败): %v", err)
}
Expand Down
33 changes: 33 additions & 0 deletions agent/embedder.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,33 @@
package agent

// Embedder 生成文本的语义向量, 用于主题相似度计算。
// 两种实现: TF-IDF (稀疏, 零依赖) 和 ONNX (稠密, 语义级)。
type Embedder interface {
// Embed 返回文本的语义向量。
Embed(text string) map[string]float64
// Name 返回嵌入器名称。
Name() string
}

// EmbedderType 嵌入器类型。
type EmbedderType string

const (
EmbedderTFIDF EmbedderType = "tfidf" // 默认: TF-IDF 稀疏向量
EmbedderONNX EmbedderType = "onnx" // ONNX Sentence Embeddings
)

// NewEmbedder 创建嵌入器实例。
// t 为类型, cacheDir 为模型缓存目录(仅 ONNX 需要)。
func NewEmbedder(t EmbedderType, cacheDir string) (Embedder, error) {
switch t {
case EmbedderONNX:
emb, err := newONNXEmbedder(cacheDir)
if err != nil {
return nil, err
}
return emb, nil
default:
return newTFIDFEmbedder(), nil
}
}
270 changes: 270 additions & 0 deletions agent/embedder_onnx.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,270 @@
package agent

import (
"fmt"
"io"
"math"
"net/http"
"os"
"path/filepath"
"strings"
"sync"
"time"

ort "github.com/getcharzp/onnxruntime_purego"
)

// onnxEmbedder 使用 ONNX Sentence Embeddings 模型生成语义向量。
// 模型: all-MiniLM-L6-v2 (384维), 首次使用从 HuggingFace 下载。
type onnxEmbedder struct {
mu sync.Mutex
session *ort.Session
vocab map[string]int32 // WordPiece 词汇表
ready bool
}

// DefaultONNXModelURL 默认 ONNX 模型下载地址(HuggingFace 镜像, 国内可用)。
const DefaultONNXModelURL = "https://hf-mirror.com/sentence-transformers/all-MiniLM-L6-v2/resolve/main/onnx/model.onnx"

// DefaultONNXVocabURL 默认词汇表下载地址(HuggingFace 镜像, 国内可用)。
const DefaultONNXVocabURL = "https://hf-mirror.com/sentence-transformers/all-MiniLM-L6-v2/resolve/main/vocab.txt"

const onnxModelFile = "embedder_model.onnx"
const onnxVocabFile = "embedder_vocab.txt"

func newONNXEmbedder(cacheDir string) (*onnxEmbedder, error) {
if err := os.MkdirAll(cacheDir, 0700); err != nil {
return nil, err
}
modelPath := filepath.Join(cacheDir, onnxModelFile)
vocabPath := filepath.Join(cacheDir, onnxVocabFile)

// 下载模型
if _, err := os.Stat(modelPath); os.IsNotExist(err) {
if err := downloadFileHTTP(DefaultONNXModelURL, modelPath); err != nil {
return nil, fmt.Errorf("下载 ONNX 模型失败: %w", err)
}
}
// 下载词汇表
if _, err := os.Stat(vocabPath); os.IsNotExist(err) {
if err := downloadFileHTTP(DefaultONNXVocabURL, vocabPath); err != nil {
return nil, fmt.Errorf("下载词汇表失败: %w", err)
}
}

// 加载词汇表
e := &onnxEmbedder{}
if err := e.loadVocab(vocabPath); err != nil {
return nil, fmt.Errorf("加载词汇表失败: %w", err)
}

// 获取共享 ONNX Runtime 引擎
engine, err := GetORTEngine(cacheDir)
if err != nil {
return nil, fmt.Errorf("ONNX Runtime 不可用: %w", err)
}

// 创建推理会话(单线程即可, 向量化很快)
session, err := engine.NewSession(modelPath, 1)
if err != nil {
return nil, fmt.Errorf("创建 ONNX 会话失败: %w", err)
}
e.session = session
e.ready = true
return e, nil
}

// loadVocab 加载 WordPiece 词汇表。
func (e *onnxEmbedder) loadVocab(path string) error {
data, err := os.ReadFile(path)
if err != nil {
return err
}
lines := strings.Split(string(data), "\n")
e.vocab = make(map[string]int32, len(lines))
for i, line := range lines {
line = strings.TrimSpace(line)
if line == "" {
continue
}
e.vocab[line] = int32(i)
}
return nil
}

// Embed 生成文本的语义向量(384维, 归一化)。
func (e *onnxEmbedder) Embed(text string) map[string]float64 {
if !e.ready || e.session == nil {
return nil
}
if text == "" {
return nil
}

e.mu.Lock()
defer e.mu.Unlock()

// 分词: 转换为 token IDs + attention mask
inputIDs, attentionMask := e.tokenize(text)
if len(inputIDs) == 0 {
return nil
}

// 构造输入张量
inputTensor, err := ort.NewTensor([]int64{1, int64(len(inputIDs))}, inputIDs)
if err != nil {
return nil
}
defer inputTensor.Destroy()

maskTensor, err := ort.NewTensor([]int64{1, int64(len(attentionMask))}, attentionMask)
if err != nil {
return nil
}
defer maskTensor.Destroy()

// ONNX 推理
outputs, err := e.session.Run(map[string]*ort.Value{
"input_ids": inputTensor,
"attention_mask": maskTensor,
})
if err != nil || len(outputs) == 0 {
return nil
}
// 获取输出(模型输出名通常为 "sentence_embedding")
var outVal *ort.Value
for _, v := range outputs {
outVal = v
break
}
defer outVal.Destroy()

// 提取向量, 384维 float32
raw, err := ort.GetTensorData[float32](outVal)
if err != nil {
return nil
}

// 转换为 float64 并归一化
vec := make(map[string]float64, 384)
var norm float64
for i, v := range raw {
vec[fmt.Sprintf("d%d", i)] = float64(v)
norm += float64(v) * float64(v)
}
norm = math.Sqrt(norm)
if norm > 0 {
for k := range vec {
vec[k] /= norm
}
}
return vec
}

func (e *onnxEmbedder) Name() string { return "onnx(all-MiniLM-L6-v2)" }

// tokenize 将文本转换为 token IDs 和 attention mask。
// 使用 WordPiece 分词算法(与 BERT 兼容)。
func (e *onnxEmbedder) tokenize(text string) ([]int64, []int64) {
const maxLen = 128
ids := make([]int64, 0, maxLen+2)
mask := make([]int64, 0, maxLen+2)

// [CLS] token (id=101 in BERT vocab)
ids = append(ids, 101)
mask = append(mask, 1)

// 对文本进行分词
runes := []rune(text)
i := 0
for i < len(runes) && len(ids) < maxLen {
r := runes[i]
if isCJK(r) {
// CJK: 每个字符单独查找词汇表
tok := strings.ToLower(string(r))
if id, ok := e.vocab[tok]; ok {
ids = append(ids, int64(id))
} else {
ids = append(ids, 100) // [UNK]
}
mask = append(mask, 1)
i++
} else if isLetterOrDigit(r) {
// 英文单词: 累积到空格或标点
var buf strings.Builder
for i < len(runes) && isLetterOrDigit(runes[i]) {
buf.WriteRune(runes[i])
i++
}
word := strings.ToLower(buf.String())
subIDs := e.wordpiece(word)
ids = append(ids, subIDs...)
for range subIDs {
mask = append(mask, 1)
}
} else {
i++
}
}

// [SEP] token (id=102)
ids = append(ids, 102)
mask = append(mask, 1)

return ids, mask
}

// wordpiece 将英文单词切分为子词。
func (e *onnxEmbedder) wordpiece(word string) []int64 {
if len(word) == 0 {
return nil
}
var ids []int64
start := 0
runes := []rune(word)
for start < len(runes) {
end := len(runes)
found := false
for end > start {
sub := string(runes[start:end])
if start > 0 {
sub = "##" + sub
}
if id, ok := e.vocab[sub]; ok {
ids = append(ids, int64(id))
start = end
found = true
break
}
end--
}
if !found {
ids = append(ids, 100) // [UNK]
start++
}
}
return ids
}

func downloadFileHTTP(url, path string) error {
resp, err := (&http.Client{Timeout: 30 * time.Second}).Get(url)
if err != nil {
return fmt.Errorf("无法下载 %s: %w", url, err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return fmt.Errorf("下载返回 %s", resp.Status)
}
tmpPath := path + ".tmp"
f, err := os.Create(tmpPath)
if err != nil {
return err
}
defer f.Close()
if _, err := io.Copy(f, io.LimitReader(resp.Body, 200<<20)); err != nil {
os.Remove(tmpPath)
return err
}
f.Close()
return os.Rename(tmpPath, path)
}
Loading