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:
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -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"))
|
||||
|
||||
Reference in New Issue
Block a user