b536672000
- PostgreSQL 切换 pgvector/pgvector:pg17 镜像;迁移 000024 建 vector 扩展、 knowledge_chunks.embedding vector(1024) + HNSW 余弦索引,retrieval_mode 放宽三态 - OllamaEmbedder 本地 bge-m3 批量嵌入,404 惰性 pull 重试,维度/超时校验,可整体关闭 - SemanticRetriever/HybridRetriever + NewRetriever 按 retrieval_mode 分发,缺 embedder 回退 FTS - 文档入库同步批量向量化;Ollama 故障降级入库 + embedding_failed 事件 - 修复 pgx CopyFrom 对 vector 列二进制编码误读:COPY 基础列后同事务 unnest 批量回填 - 修复降级路径 embeddings=nil 索引越界 panic(Add 与 Reprocess) - 知识库列表 vectorized_chunk_count + 前端三态检索模式选择与向量化覆盖率 - 单测 embedder/retrievers + 集成 TestKnowledgeVectorLifecycle 全绿 Co-Authored-By: Claude <noreply@anthropic.com>
108 lines
3.4 KiB
Go
108 lines
3.4 KiB
Go
package workbench
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"strings"
|
|
"testing"
|
|
"time"
|
|
)
|
|
|
|
// fakeOllama 模拟 Ollama 的 /api/embed 与 /api/pull。dim 固定产出向量维数,
|
|
// requireModel 为 true 时首次 /api/embed 返回 404 触发 /api/pull。
|
|
func fakeOllama(t *testing.T, dim int, requireModel bool) *httptest.Server {
|
|
t.Helper()
|
|
pulled := false
|
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
switch r.URL.Path {
|
|
case "/api/pull":
|
|
pulled = true
|
|
_ = json.NewEncoder(w).Encode(map[string]any{"status": "success"})
|
|
case "/api/embed":
|
|
if requireModel && !pulled {
|
|
w.WriteHeader(http.StatusNotFound)
|
|
_ = json.NewEncoder(w).Encode(map[string]any{"error": "model 'bge-m3' not found"})
|
|
return
|
|
}
|
|
var body struct {
|
|
Input []string `json:"input"`
|
|
}
|
|
_ = json.NewDecoder(r.Body).Decode(&body)
|
|
embeddings := make([][]float32, len(body.Input))
|
|
for i := range body.Input {
|
|
vector := make([]float32, dim)
|
|
for j := range vector {
|
|
vector[j] = float32(i + 1)
|
|
}
|
|
embeddings[i] = vector
|
|
}
|
|
_ = json.NewEncoder(w).Encode(map[string]any{"embeddings": embeddings})
|
|
default:
|
|
http.NotFound(w, r)
|
|
}
|
|
}))
|
|
t.Cleanup(server.Close)
|
|
return server
|
|
}
|
|
|
|
func TestOllamaEmbedderBatching(t *testing.T) {
|
|
server := fakeOllama(t, 4, false)
|
|
embedder := NewOllamaEmbedder(OllamaEmbedderConfig{BaseURL: server.URL, Model: "bge-m3", Dim: 4, BatchSize: 2, Timeout: 5 * time.Second})
|
|
texts := []string{"a", "b", "c", "d", "e"}
|
|
vectors, err := embedder.Embed(context.Background(), texts)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if len(vectors) != len(texts) {
|
|
t.Fatalf("got %d vectors want %d", len(vectors), len(texts))
|
|
}
|
|
for i, vector := range vectors {
|
|
if len(vector) != 4 {
|
|
t.Fatalf("vector %d has dim %d want 4", i, len(vector))
|
|
}
|
|
}
|
|
if embedder.Dim() != 4 {
|
|
t.Fatalf("Dim()=%d want 4", embedder.Dim())
|
|
}
|
|
}
|
|
|
|
func TestOllamaEmbedderDimensionMismatch(t *testing.T) {
|
|
server := fakeOllama(t, 3, false)
|
|
embedder := NewOllamaEmbedder(OllamaEmbedderConfig{BaseURL: server.URL, Model: "bge-m3", Dim: 4, BatchSize: 2, Timeout: 5 * time.Second})
|
|
_, err := embedder.Embed(context.Background(), []string{"x"})
|
|
if err == nil || !strings.Contains(err.Error(), "dimension mismatch") {
|
|
t.Fatalf("expected dimension mismatch error, got %v", err)
|
|
}
|
|
}
|
|
|
|
func TestOllamaEmbedderPullsModelOn404(t *testing.T) {
|
|
server := fakeOllama(t, 4, true)
|
|
embedder := NewOllamaEmbedder(OllamaEmbedderConfig{BaseURL: server.URL, Model: "bge-m3", Dim: 4, BatchSize: 8, Timeout: 5 * time.Second})
|
|
vectors, err := embedder.Embed(context.Background(), []string{"hello", "world"})
|
|
if err != nil {
|
|
t.Fatalf("expected pull-then-retry to succeed, got %v", err)
|
|
}
|
|
if len(vectors) != 2 {
|
|
t.Fatalf("got %d vectors want 2", len(vectors))
|
|
}
|
|
}
|
|
|
|
func TestOllamaEmbedderEmptyInput(t *testing.T) {
|
|
embedder := NewOllamaEmbedder(OllamaEmbedderConfig{BaseURL: "http://127.0.0.1:1", Model: "bge-m3", Dim: 4, BatchSize: 2, Timeout: time.Second})
|
|
vectors, err := embedder.Embed(context.Background(), nil)
|
|
if err != nil || len(vectors) != 0 {
|
|
t.Fatalf("empty input should return nil without HTTP: vectors=%d err=%v", len(vectors), err)
|
|
}
|
|
}
|
|
|
|
func TestFormatVector(t *testing.T) {
|
|
if got := formatVector([]float32{1, 2.5, -0.25}); got != "[1,2.5,-0.25]" {
|
|
t.Fatalf("formatVector got %q", got)
|
|
}
|
|
if got := formatVector([]float32{}); got != "[]" {
|
|
t.Fatalf("formatVector empty got %q", got)
|
|
}
|
|
}
|