Files
superidou b536672000 feat(m8): P2 pgvector + Ollama 向量化与语义检索
- 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>
2026-08-12 15:16:32 +08:00

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)
}
}