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>
This commit is contained in:
ben
2026-08-12 15:16:32 +08:00
parent 6708c226a5
commit b536672000
26 changed files with 966 additions and 60 deletions
+42 -1
View File
@@ -22,8 +22,9 @@ type Config struct {
Audit Audit
Outbox Outbox
RuntimeData RuntimeData
Shadow Shadow
Shadow Shadow
ObjectStorage ObjectStorage
Embeddings Embeddings
}
type Server struct {
@@ -125,6 +126,17 @@ type ObjectStorage struct {
MaxFileBytes int64
}
// Embeddings 配置本地 Ollama 向量化(默认 bge-m3)。Enabled 为 false 时网关不构造
// OllamaEmbedder,知识库退回纯 FTS 检索,AddKnowledgeDocument 不生成 embedding。
type Embeddings struct {
Enabled bool
BaseURL string
Model string
Dim int
BatchSize int
Timeout time.Duration
}
func Load() (Config, error) {
cfg := Config{
Environment: env("APP_ENV", "local"),
@@ -202,6 +214,14 @@ func Load() (Config, error) {
UseSSL: boolValue("S3_USE_SSL", false),
MaxFileBytes: int64Value("S3_MAX_FILE_BYTES", 128<<20),
},
Embeddings: Embeddings{
Enabled: boolValue("EMBEDDINGS_ENABLED", true),
BaseURL: strings.TrimRight(env("OLLAMA_BASE_URL", "http://ollama:11434"), "/"),
Model: env("EMBEDDING_MODEL", "bge-m3"),
Dim: intValue("EMBEDDING_DIM", 1024),
BatchSize: intValue("EMBEDDING_BATCH_SIZE", 64),
Timeout: duration("EMBEDDING_TIMEOUT", 120*time.Second),
},
}
return cfg, cfg.Validate()
@@ -271,6 +291,27 @@ func (c Config) Validate() error {
if c.ObjectStorage.MaxFileBytes < 1<<20 || c.ObjectStorage.MaxFileBytes > 512<<20 {
errs = append(errs, errors.New("S3_MAX_FILE_BYTES must be between 1 MiB and 512 MiB"))
}
if c.Embeddings.Enabled {
if err := validateHTTPURL(c.Embeddings.BaseURL); err != nil {
errs = append(errs, fmt.Errorf("OLLAMA_BASE_URL: %w", err))
}
if c.Embeddings.Model == "" {
errs = append(errs, errors.New("EMBEDDING_MODEL is required when embeddings are enabled"))
}
if c.Embeddings.Dim < 128 || c.Embeddings.Dim > 8192 {
errs = append(errs, errors.New("EMBEDDING_DIM must be between 128 and 8192"))
}
if c.Embeddings.Dim != 1024 {
// 知识库 embedding 列固定为 vector(1024);维度不符会让入库向量报错。
errs = append(errs, errors.New("EMBEDDING_DIM must be 1024 to match the vector(1024) column"))
}
if c.Embeddings.BatchSize < 1 || c.Embeddings.BatchSize > 512 {
errs = append(errs, errors.New("EMBEDDING_BATCH_SIZE must be between 1 and 512"))
}
if c.Embeddings.Timeout < time.Second || c.Embeddings.Timeout > 30*time.Minute {
errs = append(errs, errors.New("EMBEDDING_TIMEOUT must be between 1s and 30m"))
}
}
return errors.Join(errs...)
}
+1 -1
View File
@@ -417,7 +417,7 @@ func (h *AdminHTTPHandler) searchKnowledge(w http.ResponseWriter, r *http.Reques
if !decodeAsset(w, r, &p) {
return
}
items, err := NewPostgreSQLRetriever(h.service).Search(r.Context(), r.PathValue("id"), p.Query, p.TopK)
items, err := NewRetriever(h.service, h.service.Embedder()).Search(r.Context(), r.PathValue("id"), p.Query, p.TopK)
if err != nil {
assetError(w, err)
return
+189
View File
@@ -0,0 +1,189 @@
package workbench
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"time"
)
// Embedder 把文本批量转换成固定维度向量。知识库向量化/语义检索依赖该接口,
// Embed 失败时调用方应优雅降级(文档照常入库、embedding 置 NULL)。
type Embedder interface {
Embed(ctx context.Context, texts []string) ([][]float32, error)
Dim() int
}
// OllamaEmbedder 调用本地 Ollama 的 /api/embed(批量),默认模型 bge-m3(1024 维)。
// 模型尚未拉取时(/api/embed 返回 404)会先 POST /api/pull 拉取一次再重试。
type OllamaEmbedder struct {
baseURL string
model string
dim int
batchSize int
timeout time.Duration
client *http.Client
}
// OllamaEmbedderConfig 是 NewOllamaEmbedder 的参数;BaseURL 与 Model 已去除首尾空白。
type OllamaEmbedderConfig struct {
BaseURL string
Model string
Dim int
BatchSize int
Timeout time.Duration
}
func NewOllamaEmbedder(cfg OllamaEmbedderConfig) *OllamaEmbedder {
client := &http.Client{Timeout: cfg.Timeout}
if cfg.Dim <= 0 {
cfg.Dim = 1024
}
if cfg.BatchSize <= 0 {
cfg.BatchSize = 64
}
return &OllamaEmbedder{
baseURL: cfg.BaseURL,
model: cfg.Model,
dim: cfg.Dim,
batchSize: cfg.BatchSize,
timeout: cfg.Timeout,
client: client,
}
}
func (o *OllamaEmbedder) Dim() int { return o.dim }
type ollamaEmbedResponse struct {
Embeddings [][]float32 `json:"embeddings"`
}
type ollamaErrorResponse struct {
Error string `json:"error"`
}
// Embed 把 texts 按 batchSize 切批调用 Ollama。任何一次调用失败都会返回错误,
// 由调用方决定是否降级(知识库入库语义)。
func (o *OllamaEmbedder) Embed(ctx context.Context, texts []string) ([][]float32, error) {
if len(texts) == 0 {
return nil, nil
}
all := make([][]float32, 0, len(texts))
for start := 0; start < len(texts); start += o.batchSize {
end := start + o.batchSize
if end > len(texts) {
end = len(texts)
}
batch, err := o.embedBatch(ctx, texts[start:end])
if err != nil {
return nil, fmt.Errorf("embed batch [%d:%d]: %w", start, end, err)
}
if len(batch) != end-start {
return nil, fmt.Errorf("embed batch [%d:%d] returned %d vectors for %d texts", start, end, len(batch), end-start)
}
for _, vector := range batch {
if len(vector) != o.dim {
return nil, fmt.Errorf("embedding dimension mismatch: got %d want %d", len(vector), o.dim)
}
all = append(all, vector)
}
}
return all, nil
}
func (o *OllamaEmbedder) embedBatch(ctx context.Context, texts []string) ([][]float32, error) {
body, err := json.Marshal(map[string]any{"model": o.model, "input": texts})
if err != nil {
return nil, err
}
resp, err := o.do(ctx, "/api/embed", body)
if err != nil {
return nil, err
}
return resp, nil
}
func (o *OllamaEmbedder) do(ctx context.Context, path string, body []byte) ([][]float32, error) {
req, err := http.NewRequestWithContext(ctx, http.MethodPost, o.baseURL+path, bytes.NewReader(body))
if err != nil {
return nil, err
}
req.Header.Set("Content-Type", "application/json")
httpResp, err := o.client.Do(req)
if err != nil {
return nil, err
}
defer httpResp.Body.Close()
raw, _ := io.ReadAll(io.LimitReader(httpResp.Body, 8<<20))
if httpResp.StatusCode == http.StatusNotFound {
// 模型未拉取:先 pull(流式关闭)再重试一次;仍失败则返回明确错误。
if err := o.pullModel(ctx); err != nil {
return nil, fmt.Errorf("model %s not present and pull failed: %w", o.model, err)
}
return o.retryOnce(ctx, path, body)
}
if httpResp.StatusCode != http.StatusOK {
return nil, o.decodeError(httpResp.StatusCode, raw)
}
var parsed ollamaEmbedResponse
if err := json.Unmarshal(raw, &parsed); err != nil {
return nil, fmt.Errorf("decode embed response: %w", err)
}
return parsed.Embeddings, nil
}
func (o *OllamaEmbedder) retryOnce(ctx context.Context, path string, body []byte) ([][]float32, error) {
req, err := http.NewRequestWithContext(ctx, http.MethodPost, o.baseURL+path, bytes.NewReader(body))
if err != nil {
return nil, err
}
req.Header.Set("Content-Type", "application/json")
httpResp, err := o.client.Do(req)
if err != nil {
return nil, err
}
defer httpResp.Body.Close()
raw, _ := io.ReadAll(io.LimitReader(httpResp.Body, 8<<20))
if httpResp.StatusCode != http.StatusOK {
return nil, o.decodeError(httpResp.StatusCode, raw)
}
var parsed ollamaEmbedResponse
if err := json.Unmarshal(raw, &parsed); err != nil {
return nil, fmt.Errorf("decode embed response after pull: %w", err)
}
return parsed.Embeddings, nil
}
func (o *OllamaEmbedder) pullModel(ctx context.Context) error {
body, err := json.Marshal(map[string]any{"model": o.model, "stream": false})
if err != nil {
return err
}
req, err := http.NewRequestWithContext(ctx, http.MethodPost, o.baseURL+"/api/pull", bytes.NewReader(body))
if err != nil {
return err
}
req.Header.Set("Content-Type", "application/json")
resp, err := o.client.Do(req)
if err != nil {
return err
}
defer resp.Body.Close()
raw, _ := io.ReadAll(io.LimitReader(resp.Body, 8<<20))
if resp.StatusCode != http.StatusOK {
return o.decodeError(resp.StatusCode, raw)
}
return nil
}
func (o *OllamaEmbedder) decodeError(status int, raw []byte) error {
var parsed ollamaErrorResponse
if err := json.Unmarshal(raw, &parsed); err == nil && parsed.Error != "" {
return errors.New(parsed.Error)
}
return fmt.Errorf("ollama request failed with status %d", status)
}
+107
View File
@@ -0,0 +1,107 @@
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)
}
}
+74 -22
View File
@@ -5,6 +5,7 @@ import (
"crypto/sha256"
"encoding/hex"
"errors"
"log/slog"
"regexp"
"sort"
"strings"
@@ -24,23 +25,11 @@ type PostgreSQLRetriever struct {
}
}
func NewPostgreSQLRetriever(service *Service) *PostgreSQLRetriever {
return &PostgreSQLRetriever{pool: service.pool}
}
func (r *PostgreSQLRetriever) Search(ctx context.Context, knowledgeBaseID, query string, topK int) ([]SearchHit, error) {
query = strings.TrimSpace(query)
if query == "" {
return nil, errors.New("检索词不能为空")
}
if len(query) > 6000 {
return nil, errors.New("检索词过长")
}
if topK < 1 {
topK = 4
}
if topK > 20 {
topK = 20
var err error
query, topK, err = validateSearch(query, topK)
if err != nil {
return nil, err
}
tokens := searchTokens(query)
rows, err := r.pool.Query(ctx, `WITH q AS (SELECT plainto_tsquery('simple',$2) AS tsq,lower($2) AS raw), tokens AS (SELECT unnest($4::text[]) AS token)
@@ -114,8 +103,11 @@ func validateKnowledgeBase(k *KnowledgeBase) error {
if k.RetrievalMode == "" {
k.RetrievalMode = "postgres_fts"
}
if k.RetrievalMode != "postgres_fts" {
return errors.New("基线仅支持 postgres_fts 检索器")
switch k.RetrievalMode {
case "postgres_fts", "vector", "hybrid":
// M8 P2:三态检索模式。vector/hybrid 需要向量化器,未启用时检索自动回退 FTS。
default:
return errors.New("retrieval_mode 仅支持 postgres_fts / vector / hybrid")
}
if k.ChunkSize == 0 {
k.ChunkSize = 800
@@ -132,12 +124,12 @@ func validateKnowledgeBase(k *KnowledgeBase) error {
}
const knowledgeBaseSelect = `SELECT k.id::text,k.name,k.description,k.retrieval_mode,k.chunk_size,k.chunk_overlap,k.department_ids::text[],k.enabled,k.revision,
count(DISTINCT d.id)::int,count(c.id)::int,k.created_at,k.updated_at
count(DISTINCT d.id)::int,count(c.id)::int,count(c.embedding)::int,k.created_at,k.updated_at
FROM gateway.knowledge_bases k LEFT JOIN gateway.knowledge_documents d ON d.knowledge_base_id=k.id LEFT JOIN gateway.knowledge_chunks c ON c.document_id=d.id`
func scanKnowledgeBase(row pgx.Row) (KnowledgeBase, error) {
var k KnowledgeBase
err := row.Scan(&k.ID, &k.Name, &k.Description, &k.RetrievalMode, &k.ChunkSize, &k.ChunkOverlap, &k.DepartmentIDs, &k.Enabled, &k.Revision, &k.DocumentCount, &k.ChunkCount, &k.CreatedAt, &k.UpdatedAt)
err := row.Scan(&k.ID, &k.Name, &k.Description, &k.RetrievalMode, &k.ChunkSize, &k.ChunkOverlap, &k.DepartmentIDs, &k.Enabled, &k.Revision, &k.DocumentCount, &k.ChunkCount, &k.VectorizedChunkCount, &k.CreatedAt, &k.UpdatedAt)
return k, mapNotFound(err)
}
@@ -295,6 +287,31 @@ func ChunkText(text string, size, overlap int) []string {
return chunks
}
// embedResult 描述文档入库时的向量化结果。
type embedResult struct {
embeddings [][]float32 // 与 chunks 一一对应;needed=false 或 degraded=true 时为 nil
needed bool // 该知识库模式需要 embedding
degraded bool // 需要但计算失败,文档降级为纯 FTS 入库
}
// computeEmbeddings 在 KB 为 vector/hybrid 且已注入 embedder 时同步批量计算分块向量。
// 失败时不阻断入库:返回 degraded=true,由调用方在事务内补发 embedding_failed 事件。
func (s *Service) computeEmbeddings(ctx context.Context, kb KnowledgeBase, chunks []string) embedResult {
if s.embedder == nil || (kb.RetrievalMode != "vector" && kb.RetrievalMode != "hybrid") {
return embedResult{}
}
vectors, err := s.embedder.Embed(ctx, chunks)
if err != nil {
slog.Warn("knowledge embedding failed, storing document without vectors", "knowledge_base_id", kb.ID, "error", err)
return embedResult{needed: true, degraded: true}
}
if len(vectors) != len(chunks) {
slog.Warn("knowledge embedding count mismatch, storing document without vectors", "knowledge_base_id", kb.ID, "got", len(vectors), "want", len(chunks))
return embedResult{needed: true, degraded: true}
}
return embedResult{embeddings: vectors, needed: true}
}
func (s *Service) AddKnowledgeDocument(ctx context.Context, kbID, title, sourceType, sourceURI, content, actorID string) (KnowledgeDocument, error) {
title = strings.TrimSpace(title)
sourceType = strings.TrimSpace(sourceType)
@@ -336,22 +353,45 @@ func (s *Service) AddKnowledgeDocument(ctx context.Context, kbID, title, sourceT
if err != nil {
return KnowledgeDocument{}, err
}
embed := s.computeEmbeddings(ctx, kb, chunks)
chunkRows := make([][]any, 0, len(chunks))
chunkIDs := make([]string, 0, len(chunks))
embeddings := make([]string, 0, len(chunks))
for index, chunk := range chunks {
chunkID, idErr := newUUID()
if idErr != nil {
return KnowledgeDocument{}, idErr
}
chunkRows = append(chunkRows, []any{chunkID, kbID, docID, index, chunk})
// embed.embeddings 仅在成功算出向量时非 nil;降级(失败)时跳过回填。
if embed.embeddings != nil {
chunkIDs = append(chunkIDs, chunkID)
embeddings = append(embeddings, formatVector(embed.embeddings[index]))
}
}
// Bulk-copy all chunks in one statement instead of one INSERT per chunk;
// a 2 MiB document can split into thousands of chunks.
if _, err = tx.CopyFrom(ctx, pgx.Identifier{"gateway", "knowledge_chunks"}, []string{"id", "knowledge_base_id", "document_id", "chunk_index", "content"}, pgx.CopyFromRows(chunkRows)); err != nil {
return KnowledgeDocument{}, err
}
// pgx CopyFrom 对未知 OID(vector)列走二进制编码,直接随 COPY 传字面量会报
// "vector cannot have more than 16000 dimensions";因此向量在 COPY 之后用
// unnest 批量回填(同一事务内,单条 UPDATE)。
if embed.embeddings != nil {
// 降级时 embeddings 为空,跳过向量回填;否则空数组 UPDATE 也是无害空操作。
if _, err = tx.Exec(ctx, `UPDATE gateway.knowledge_chunks c SET embedding = v.vec::vector
FROM unnest($1::uuid[], $2::text[]) AS v(id, vec) WHERE c.id = v.id`, chunkIDs, embeddings); err != nil {
return KnowledgeDocument{}, err
}
}
if err = emit(ctx, tx, "knowledge_document.ready", "knowledge_document", docID, actorID, map[string]any{"knowledge_base_id": kbID, "chunk_count": len(chunks)}); err != nil {
return KnowledgeDocument{}, err
}
if embed.degraded {
if err = emit(ctx, tx, "knowledge_document.embedding_failed", "knowledge_document", docID, actorID, map[string]any{"knowledge_base_id": kbID, "chunk_count": len(chunks)}); err != nil {
return KnowledgeDocument{}, err
}
}
if err = tx.Commit(ctx); err != nil {
return KnowledgeDocument{}, err
}
@@ -420,13 +460,20 @@ func (s *Service) ReprocessKnowledgeDocument(ctx context.Context, kbID, id, acto
if _, err = tx.Exec(ctx, `DELETE FROM gateway.knowledge_chunks WHERE document_id=$1`, id); err != nil {
return KnowledgeDocument{}, err
}
embed := s.computeEmbeddings(ctx, kb, chunks)
for index, chunk := range chunks {
chunkID, idErr := newUUID()
if idErr != nil {
return KnowledgeDocument{}, idErr
}
if _, err = tx.Exec(ctx, `INSERT INTO gateway.knowledge_chunks(id,knowledge_base_id,document_id,chunk_index,content) VALUES($1,$2,$3,$4,$5)`, chunkID, kbID, id, index, chunk); err != nil {
return KnowledgeDocument{}, err
if embed.embeddings != nil {
if _, err = tx.Exec(ctx, `INSERT INTO gateway.knowledge_chunks(id,knowledge_base_id,document_id,chunk_index,content,embedding) VALUES($1,$2,$3,$4,$5,$6::vector)`, chunkID, kbID, id, index, chunk, formatVector(embed.embeddings[index])); err != nil {
return KnowledgeDocument{}, err
}
} else {
if _, err = tx.Exec(ctx, `INSERT INTO gateway.knowledge_chunks(id,knowledge_base_id,document_id,chunk_index,content) VALUES($1,$2,$3,$4,$5)`, chunkID, kbID, id, index, chunk); err != nil {
return KnowledgeDocument{}, err
}
}
}
err = tx.QueryRow(ctx, `UPDATE gateway.knowledge_documents SET status='ready',status_message='',chunk_count=$3,updated_at=clock_timestamp() WHERE id=$1 AND knowledge_base_id=$2 RETURNING updated_at`, id, kbID, len(chunks)).Scan(&doc.UpdatedAt)
@@ -437,6 +484,11 @@ func (s *Service) ReprocessKnowledgeDocument(ctx context.Context, kbID, id, acto
if err = emit(ctx, tx, "knowledge_document.reprocessed", "knowledge_document", id, actorID, map[string]any{"knowledge_base_id": kbID, "chunk_count": len(chunks)}); err != nil {
return KnowledgeDocument{}, err
}
if embed.degraded {
if err = emit(ctx, tx, "knowledge_document.embedding_failed", "knowledge_document", id, actorID, map[string]any{"knowledge_base_id": kbID, "chunk_count": len(chunks)}); err != nil {
return KnowledgeDocument{}, err
}
}
if err = tx.Commit(ctx); err != nil {
return KnowledgeDocument{}, err
}
+159
View File
@@ -0,0 +1,159 @@
package workbench
import (
"context"
"errors"
"sort"
"strconv"
"strings"
"github.com/jackc/pgx/v5"
)
// NewRetriever 按知识库的 retrieval_mode 分发检索器:
// - embedder 为 nil 时恒为 postgres_fts(向量化未启用)
// - 否则按 knowledge_bases.retrieval_mode 每次 Search 动态分发
// (postgres_fts / vector / hybrid),管理员修改模式后无需重启即生效。
func NewRetriever(service *Service, embedder Embedder) Retriever {
return &modeRetriever{pool: service.pool, embedder: embedder}
}
type modeRetriever struct {
pool interface {
Query(context.Context, string, ...any) (pgx.Rows, error)
QueryRow(context.Context, string, ...any) pgx.Row
}
embedder Embedder
}
func (r *modeRetriever) Search(ctx context.Context, knowledgeBaseID, query string, topK int) ([]SearchHit, error) {
if r.embedder == nil {
return (&PostgreSQLRetriever{pool: r.pool}).Search(ctx, knowledgeBaseID, query, topK)
}
var mode string
err := r.pool.QueryRow(ctx, `SELECT retrieval_mode FROM gateway.knowledge_bases WHERE id=$1`, knowledgeBaseID).Scan(&mode)
if err != nil {
return nil, mapNotFound(err)
}
switch mode {
case "vector":
return (&SemanticRetriever{pool: r.pool, embedder: r.embedder}).Search(ctx, knowledgeBaseID, query, topK)
case "hybrid":
return (&HybridRetriever{pool: r.pool, embedder: r.embedder}).Search(ctx, knowledgeBaseID, query, topK)
default:
return (&PostgreSQLRetriever{pool: r.pool}).Search(ctx, knowledgeBaseID, query, topK)
}
}
// SemanticRetriever 用 pgvector 余弦距离召回最近分块。
type SemanticRetriever struct {
pool interface {
Query(context.Context, string, ...any) (pgx.Rows, error)
}
embedder Embedder
}
func (r *SemanticRetriever) Search(ctx context.Context, knowledgeBaseID, query string, topK int) ([]SearchHit, error) {
query, topK, err := validateSearch(query, topK)
if err != nil {
return nil, err
}
vectors, err := r.embedder.Embed(ctx, []string{query})
if err != nil {
return nil, err
}
if len(vectors) == 0 {
return nil, errors.New("查询向量为空")
}
rows, err := r.pool.Query(ctx, `SELECT c.id::text,c.document_id::text,d.title,c.chunk_index,c.content,
(1-(c.embedding <=> $2::vector))::float8 AS score
FROM gateway.knowledge_chunks c JOIN gateway.knowledge_documents d ON d.id=c.document_id JOIN gateway.knowledge_bases k ON k.id=c.knowledge_base_id
WHERE c.knowledge_base_id=$1 AND k.enabled AND d.status='ready' AND c.embedding IS NOT NULL
ORDER BY c.embedding <=> $2::vector LIMIT $3`, knowledgeBaseID, formatVector(vectors[0]), topK)
if err != nil {
return nil, err
}
defer rows.Close()
hits := []SearchHit{}
for rows.Next() {
var h SearchHit
if err = rows.Scan(&h.ChunkID, &h.DocumentID, &h.DocumentTitle, &h.ChunkIndex, &h.Content, &h.Score); err != nil {
return nil, err
}
hits = append(hits, h)
}
return hits, rows.Err()
}
// HybridRetriever 融合 FTS 与语义召回:分别取 topK*2,按 chunk_id 去重取高分,
// 再截断到 topK。任一子检索失败时回退到另一路,保证向量化降级仍有结果。
type HybridRetriever struct {
pool interface {
Query(context.Context, string, ...any) (pgx.Rows, error)
}
embedder Embedder
}
func (r *HybridRetriever) Search(ctx context.Context, knowledgeBaseID, query string, topK int) ([]SearchHit, error) {
_, topK, err := validateSearch(query, topK)
if err != nil {
return nil, err
}
fts, ftsErr := (&PostgreSQLRetriever{pool: r.pool}).Search(ctx, knowledgeBaseID, query, topK*2)
sem, semErr := (&SemanticRetriever{pool: r.pool, embedder: r.embedder}).Search(ctx, knowledgeBaseID, query, topK*2)
if ftsErr != nil && semErr != nil {
return nil, ftsErr
}
best := make(map[string]SearchHit)
for _, h := range fts {
if existing, ok := best[h.ChunkID]; !ok || h.Score > existing.Score {
best[h.ChunkID] = h
}
}
for _, h := range sem {
if existing, ok := best[h.ChunkID]; !ok || h.Score > existing.Score {
best[h.ChunkID] = h
}
}
merged := make([]SearchHit, 0, len(best))
for _, h := range best {
merged = append(merged, h)
}
sort.Slice(merged, func(i, j int) bool { return merged[i].Score > merged[j].Score })
if len(merged) > topK {
merged = merged[:topK]
}
return merged, nil
}
func validateSearch(query string, topK int) (string, int, error) {
query = strings.TrimSpace(query)
if query == "" {
return "", 0, errors.New("检索词不能为空")
}
if len(query) > 6000 {
return "", 0, errors.New("检索词过长")
}
if topK < 1 {
topK = 4
}
if topK > 20 {
topK = 20
}
return query, topK, nil
}
// formatVector 把 []float32 转成 pgvector 字面量字符串 "[0.1,0.2,...]",
// 直接以 $n::vector 参数传入,避免引入 pgvector-go 依赖。
func formatVector(vector []float32) string {
var builder strings.Builder
builder.WriteByte('[')
for index, value := range vector {
if index > 0 {
builder.WriteByte(',')
}
builder.WriteString(strconv.FormatFloat(float64(value), 'g', -1, 32))
}
builder.WriteByte(']')
return builder.String()
}
+108
View File
@@ -0,0 +1,108 @@
package workbench
import (
"context"
"errors"
"testing"
)
func TestValidateKnowledgeBaseRetrievalModes(t *testing.T) {
for _, mode := range []string{"postgres_fts", "vector", "hybrid"} {
kb := KnowledgeBase{Name: "kb", RetrievalMode: mode, ChunkSize: 800, ChunkOverlap: 100}
if err := validateKnowledgeBase(&kb); err != nil {
t.Fatalf("mode %s should be accepted, got %v", mode, err)
}
}
kb := KnowledgeBase{Name: "kb", RetrievalMode: "bm25", ChunkSize: 800, ChunkOverlap: 100}
if err := validateKnowledgeBase(&kb); err == nil {
t.Fatal("unknown retrieval mode should be rejected")
}
// 空模式回填默认值
empty := KnowledgeBase{Name: "kb", ChunkSize: 800, ChunkOverlap: 100}
if err := validateKnowledgeBase(&empty); err != nil || empty.RetrievalMode != "postgres_fts" {
t.Fatalf("empty mode should default to postgres_fts, got %q err=%v", empty.RetrievalMode, err)
}
}
func TestValidateSearchClamps(t *testing.T) {
if _, topK, err := validateSearch(" 关键词 ", 0); err != nil || topK != 4 {
t.Fatalf("empty topK should clamp to 4, got %d err=%v", topK, err)
}
if _, topK, err := validateSearch("关键词", 99); err != nil || topK != 20 {
t.Fatalf("topK should clamp to 20, got %d err=%v", topK, err)
}
if _, _, err := validateSearch(" ", 4); err == nil {
t.Fatal("blank query should be rejected")
}
long := make([]byte, 6001)
for i := range long {
long[i] = 'a'
}
if _, _, err := validateSearch(string(long), 4); err == nil {
t.Fatal("over-long query should be rejected")
}
}
// failingEmbedder 总是报错,用于验证 computeEmbeddings 优雅降级。
type failingEmbedder struct{}
func (failingEmbedder) Embed(context.Context, []string) ([][]float32, error) { return nil, errors.New("embedding service down") }
func (failingEmbedder) Dim() int { return 1024 }
func TestComputeEmbeddingsDispatch(t *testing.T) {
service := NewService(nil)
chunks := []string{"a", "b"}
// 未注入 embedder:任何模式都不算向量。
if result := service.computeEmbeddings(context.Background(), KnowledgeBase{ID: "kb", RetrievalMode: "vector"}, chunks); result.needed {
t.Fatal("no embedder should not mark embedding needed")
}
// 已注入 embedder 但模式为 postgres_fts:不算向量。
service.SetEmbedder(fakeGoodEmbedder{})
if result := service.computeEmbeddings(context.Background(), KnowledgeBase{ID: "kb", RetrievalMode: "postgres_fts"}, chunks); result.needed {
t.Fatal("postgres_fts mode should not embed")
}
// vector 模式 + 失败 embedder:needed + degraded,embedding 为 nil。
service.SetEmbedder(failingEmbedder{})
result := service.computeEmbeddings(context.Background(), KnowledgeBase{ID: "kb", RetrievalMode: "hybrid"}, chunks)
if !result.needed || !result.degraded || result.embeddings != nil {
t.Fatalf("expected needed+degraded with nil embeddings, got %+v", result)
}
// vector 模式 + 正常 embedder:返回与 chunks 等长的向量。
service.SetEmbedder(fakeGoodEmbedder{})
result = service.computeEmbeddings(context.Background(), KnowledgeBase{ID: "kb", RetrievalMode: "vector"}, chunks)
if !result.needed || result.degraded || len(result.embeddings) != len(chunks) {
t.Fatalf("expected aligned embeddings, got %+v", result)
}
}
type fakeGoodEmbedder struct{}
func (fakeGoodEmbedder) Embed(_ context.Context, texts []string) ([][]float32, error) {
vectors := make([][]float32, len(texts))
for i := range texts {
vectors[i] = make([]float32, 4)
for j := range vectors[i] {
vectors[i][j] = float32(i + 1)
}
}
return vectors, nil
}
func (fakeGoodEmbedder) Dim() int { return 4 }
func TestNewRetrieverNilEmbedder(t *testing.T) {
// embedder 为 nil 时 NewRetriever 返回的检索器直接走 FTS 分支,不会查询检索模式。
service := NewService(nil)
retriever := NewRetriever(service, nil)
if retriever == nil {
t.Fatal("NewRetriever with nil embedder should not return nil")
}
// FTS 分支在无连接池时只在 query 校验处返回,不 panic。
if _, err := retriever.Search(context.Background(), "kb", " ", 4); err == nil {
t.Fatal("blank query should fail at validation before touching the pool")
}
}
+10 -1
View File
@@ -12,10 +12,19 @@ import (
"github.com/jackc/pgx/v5/pgxpool"
)
type Service struct{ pool *pgxpool.Pool }
type Service struct {
pool *pgxpool.Pool
embedder Embedder
}
func NewService(pool *pgxpool.Pool) *Service { return &Service{pool: pool} }
// SetEmbedder 注入向量化器(M8 P2)。为 nil 时知识库退回纯 FTS,不生成 embedding。
func (s *Service) SetEmbedder(embedder Embedder) { s.embedder = embedder }
// Embedder 返回当前向量化器(可能为 nil)。
func (s *Service) Embedder() Embedder { return s.embedder }
func newUUID() (string, error) { return platformid.NewUUID() }
func emit(ctx context.Context, tx pgx.Tx, eventType, aggregateType, aggregateID, actorID string, values map[string]any) error {
+6 -5
View File
@@ -82,11 +82,12 @@ type KnowledgeBase struct {
ChunkOverlap int `json:"chunk_overlap"`
DepartmentIDs []string `json:"department_ids"`
Enabled bool `json:"enabled"`
Revision int64 `json:"revision"`
DocumentCount int `json:"document_count"`
ChunkCount int `json:"chunk_count"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
Revision int64 `json:"revision"`
DocumentCount int `json:"document_count"`
ChunkCount int `json:"chunk_count"`
VectorizedChunkCount int `json:"vectorized_chunk_count"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
type KnowledgeDocument struct {
@@ -0,0 +1,90 @@
package workbench
import (
"context"
"os"
"strings"
"testing"
"time"
"aigateway.local/core/internal/platform/config"
"aigateway.local/core/internal/platform/database"
)
// TestKnowledgeVectorLifecycle 跑真实 PostgreSQL(pgvector)+ Ollama:
// vector 模式知识库导入文档 → embedding 非空 → SemanticRetriever 语义命中;
// 并验证 embedder 失败时文档照常入库且发出 embedding_failed 事件。
// 需要 WORKBENCH_TEST_DATABASE_URL 与 WORKBENCH_TEST_OLLAMA_URL。
func TestKnowledgeVectorLifecycle(t *testing.T) {
databaseURL := os.Getenv("WORKBENCH_TEST_DATABASE_URL")
ollamaURL := os.Getenv("WORKBENCH_TEST_OLLAMA_URL")
if databaseURL == "" || ollamaURL == "" {
t.Skip("WORKBENCH_TEST_DATABASE_URL and WORKBENCH_TEST_OLLAMA_URL are not set")
}
ctx := context.Background()
pool, err := database.Open(ctx, config.Database{URL: databaseURL, MaxConns: 8, MinConns: 0})
if err != nil {
t.Fatal(err)
}
defer pool.Close()
actorID := "55555555-5555-4555-8555-555555555555"
_, err = pool.Exec(ctx, `INSERT INTO gateway.admin_accounts(id,username,password_hash,role) VALUES($1,'m8-vector-test','test','superadmin') ON CONFLICT(id) DO NOTHING`, actorID)
if err != nil {
t.Fatal(err)
}
name := "m8-vector-integration-kb"
cleanup := func() {
_, _ = pool.Exec(ctx, `DELETE FROM gateway.knowledge_bases WHERE name=$1 OR name=$2`, name, name+"-fts")
}
cleanup()
defer cleanup()
assets := NewService(pool)
assets.SetEmbedder(NewOllamaEmbedder(OllamaEmbedderConfig{BaseURL: ollamaURL, Model: "bge-m3", Dim: 1024, BatchSize: 8, Timeout: 60 * time.Second}))
kb, err := assets.SaveKnowledgeBase(ctx, KnowledgeBase{Name: name, Description: "vector itest", RetrievalMode: "vector", ChunkSize: 300, ChunkOverlap: 40, DepartmentIDs: []string{}, Enabled: true}, actorID, true)
if err != nil {
t.Fatal(err)
}
doc, err := assets.AddKnowledgeDocument(ctx, kb.ID, "向量检索测试", "text", "", "pgvector 语义检索依赖 bge-m3 向量。Ollama 本地生成嵌入。", actorID)
if err != nil {
t.Fatal(err)
}
var embedded int
if err = pool.QueryRow(ctx, `SELECT count(*) FROM gateway.knowledge_chunks WHERE document_id=$1 AND embedding IS NOT NULL`, doc.ID).Scan(&embedded); err != nil {
t.Fatal(err)
}
if embedded == 0 || embedded != doc.ChunkCount {
t.Fatalf("expected all chunks embedded, got %d/%d", embedded, doc.ChunkCount)
}
// 语义检索:查询词与正文无字面重合也应命中(余弦相似度)。
hitChunks, err := (&SemanticRetriever{pool: pool, embedder: assets.Embedder()}).Search(ctx, kb.ID, "语义相似度匹配", 4)
if err != nil {
t.Fatalf("semantic search: %v", err)
}
if len(hitChunks) == 0 {
t.Fatal("expected semantic hit")
}
if !strings.Contains(hitChunks[0].Content, "pgvector") {
t.Fatalf("expected pgvector content in top hit, got %q", hitChunks[0].Content)
}
// 降级路径:embedder 失败 → 文档照常入库 + embedding_failed 事件。
assets.SetEmbedder(failingEmbedder{})
kbFts, err := assets.SaveKnowledgeBase(ctx, KnowledgeBase{Name: name + "-fts", Description: "degraded", RetrievalMode: "vector", ChunkSize: 300, ChunkOverlap: 40, DepartmentIDs: []string{}, Enabled: true}, actorID, true)
if err != nil {
t.Fatal(err)
}
degradedDoc, err := assets.AddKnowledgeDocument(ctx, kbFts.ID, "降级测试", "text", "", "没有向量也能入库。", actorID)
if err != nil {
t.Fatalf("degraded add should succeed, got %v", err)
}
var failed bool
if err = pool.QueryRow(ctx, `SELECT EXISTS(SELECT 1 FROM gateway.outbox_events WHERE event_type='knowledge_document.embedding_failed' AND aggregate_id=$1)`, degradedDoc.ID).Scan(&failed); err != nil {
t.Fatal(err)
}
if !failed {
t.Fatal("expected embedding_failed event when embedder is down")
}
}
@@ -66,7 +66,7 @@ func TestWorkbenchPostgreSQLLifecycle(t *testing.T) {
if doc.ChunkCount == 0 {
t.Fatal("expected chunks")
}
hits, err := NewPostgreSQLRetriever(assets).Search(ctx, kb.ID, "不可变运行时快照", 4)
hits, err := NewRetriever(assets, nil).Search(ctx, kb.ID, "不可变运行时快照", 4)
if err != nil || len(hits) == 0 {
t.Fatalf("hits=%#v err=%v", hits, err)
}
@@ -116,7 +116,7 @@ func TestWorkbenchPostgreSQLLifecycle(t *testing.T) {
}
_ = json.NewEncoder(w).Encode(map[string]any{"choices": []any{map[string]any{"message": map[string]any{"role": "assistant", "content": "完成"}}}})
})
runtime := NewRuntimeHTTPHandler(assets, tools, NewPostgreSQLRetriever(assets), staticPrincipalAuthenticator{}, fakeGateway, MarketplaceDeps{})
runtime := NewRuntimeHTTPHandler(assets, tools, NewRetriever(assets, nil), staticPrincipalAuthenticator{}, fakeGateway, MarketplaceDeps{})
runtimeRequest := httptest.NewRequest(http.MethodPost, "/v1/applications/m4_app/chat/completions", bytes.NewBufferString(`{"messages":[{"role":"user","content":"不可变运行时快照是什么?"}],"variables":{"question":"架构"}}`))
runtimeRequest.Header.Set("Authorization", "Bearer test")
runtimeRequest = runtimeRequest.WithContext(gateway.WithRequestID(runtimeRequest.Context(), "m4-runtime"))