mirror of
https://github.com/Tencent/WeKnora.git
synced 2026-09-24 16:29:01 +08:00
feat: add VectorStore service layer with types and connection testing
Add VectorStoreService with CRUD validation, duplicate checking (DB + env stores), and engine-specific health checks for 6 vector database types. Include VectorStoreResponse DTO, env store builder, engine type metadata, and comprehensive unit tests.
This commit is contained in:
@@ -59,6 +59,15 @@ func (r *vectorStoreRepository) Update(ctx context.Context, store *types.VectorS
|
||||
).Select("name").Updates(store).Error
|
||||
}
|
||||
|
||||
// UpdateConnectionConfig updates only the connection_config JSONB column.
|
||||
// Used for saving auto-detected metadata (e.g., server version) without
|
||||
// touching user-immutable fields like engine_type or index_config.
|
||||
func (r *vectorStoreRepository) UpdateConnectionConfig(ctx context.Context, store *types.VectorStore) error {
|
||||
return r.db.WithContext(ctx).Model(&types.VectorStore{}).Where(
|
||||
"id = ? AND tenant_id = ?", store.ID, store.TenantID,
|
||||
).Select("connection_config").Updates(store).Error
|
||||
}
|
||||
|
||||
// Delete soft-deletes a vector store
|
||||
func (r *vectorStoreRepository) Delete(ctx context.Context, tenantID uint64, id string) error {
|
||||
return r.db.WithContext(ctx).Where(
|
||||
|
||||
@@ -0,0 +1,121 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
|
||||
"github.com/Tencent/WeKnora/internal/errors"
|
||||
"github.com/Tencent/WeKnora/internal/logger"
|
||||
"github.com/Tencent/WeKnora/internal/types"
|
||||
"github.com/Tencent/WeKnora/internal/types/interfaces"
|
||||
secutils "github.com/Tencent/WeKnora/internal/utils"
|
||||
)
|
||||
|
||||
// vectorStoreService implements interfaces.VectorStoreService
|
||||
type vectorStoreService struct {
|
||||
repo interfaces.VectorStoreRepository
|
||||
}
|
||||
|
||||
// NewVectorStoreService creates a new vector store service
|
||||
func NewVectorStoreService(
|
||||
repo interfaces.VectorStoreRepository,
|
||||
) interfaces.VectorStoreService {
|
||||
return &vectorStoreService{repo: repo}
|
||||
}
|
||||
|
||||
// CreateStore validates and creates a new vector store.
|
||||
func (s *vectorStoreService) CreateStore(ctx context.Context, store *types.VectorStore) error {
|
||||
// 1. Basic validation (name, engine_type, tenant_id)
|
||||
if err := store.Validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 2. Engine-specific connection config validation
|
||||
if err := validateConnectionConfig(store.EngineType, store.ConnectionConfig); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 3. Duplicate check — DB stores
|
||||
endpoint := store.ConnectionConfig.GetEndpoint()
|
||||
indexName := store.IndexConfig.GetIndexNameOrDefault(store.EngineType)
|
||||
|
||||
exists, err := s.repo.ExistsByEndpointAndIndex(ctx, store.TenantID, store.EngineType, endpoint, indexName)
|
||||
if err != nil {
|
||||
return errors.NewInternalServerError("failed to check for duplicates")
|
||||
}
|
||||
if exists {
|
||||
return errors.NewConflictError("a vector store with the same endpoint and index already exists")
|
||||
}
|
||||
|
||||
// 4. Duplicate check — env stores (pure function, no os.Getenv in types)
|
||||
for _, envStore := range types.BuildEnvVectorStores(os.Getenv("RETRIEVE_DRIVER"), os.Getenv) {
|
||||
if envStore.EngineType == store.EngineType &&
|
||||
envStore.ConnectionConfig.GetEndpoint() == endpoint &&
|
||||
envStore.IndexConfig.GetIndexNameOrDefault(store.EngineType) == indexName {
|
||||
return errors.NewConflictError(
|
||||
"a vector store with the same endpoint and index is already configured via environment variables")
|
||||
}
|
||||
}
|
||||
|
||||
// 5. Persist
|
||||
logger.Infof(ctx, "Creating vector store: tenant=%d, name=%s, engine=%s",
|
||||
store.TenantID, secutils.SanitizeForLog(store.Name), store.EngineType)
|
||||
return s.repo.Create(ctx, store)
|
||||
}
|
||||
|
||||
// UpdateStore updates an existing vector store (name only).
|
||||
func (s *vectorStoreService) UpdateStore(ctx context.Context, store *types.VectorStore) error {
|
||||
if store.TenantID == 0 {
|
||||
return errors.NewValidationError("tenant_id is required")
|
||||
}
|
||||
if store.Name == "" {
|
||||
return errors.NewValidationError("name is required")
|
||||
}
|
||||
|
||||
logger.Infof(ctx, "Updating vector store: tenant=%d, id=%s", store.TenantID, store.ID)
|
||||
return s.repo.Update(ctx, store)
|
||||
}
|
||||
|
||||
// DeleteStore deletes a vector store by tenant + id.
|
||||
// Phase 2: KB binding check will be added here.
|
||||
func (s *vectorStoreService) DeleteStore(ctx context.Context, tenantID uint64, id string) error {
|
||||
logger.Infof(ctx, "Deleting vector store: tenant=%d, id=%s", tenantID, id)
|
||||
return s.repo.Delete(ctx, tenantID, id)
|
||||
}
|
||||
|
||||
// SaveDetectedVersion updates the connection_config.version for a stored vector store.
|
||||
// Works on a copy to avoid mutating the caller's object.
|
||||
func (s *vectorStoreService) SaveDetectedVersion(ctx context.Context, store *types.VectorStore, version string) error {
|
||||
updated := *store
|
||||
updated.ConnectionConfig.Version = version
|
||||
return s.repo.UpdateConnectionConfig(ctx, &updated)
|
||||
}
|
||||
|
||||
// validateConnectionConfig validates required fields per engine type.
|
||||
func validateConnectionConfig(engineType types.RetrieverEngineType, config types.ConnectionConfig) error {
|
||||
switch engineType {
|
||||
case types.ElasticsearchRetrieverEngineType:
|
||||
if config.Addr == "" {
|
||||
return errors.NewValidationError("addr is required for elasticsearch")
|
||||
}
|
||||
case types.PostgresRetrieverEngineType:
|
||||
if !config.UseDefaultConnection && config.Addr == "" {
|
||||
return errors.NewValidationError("addr or use_default_connection is required for postgres")
|
||||
}
|
||||
case types.QdrantRetrieverEngineType:
|
||||
if config.Host == "" {
|
||||
return errors.NewValidationError("host is required for qdrant")
|
||||
}
|
||||
case types.MilvusRetrieverEngineType:
|
||||
if config.Addr == "" {
|
||||
return errors.NewValidationError("addr is required for milvus")
|
||||
}
|
||||
case types.WeaviateRetrieverEngineType:
|
||||
if config.Host == "" {
|
||||
return errors.NewValidationError("host is required for weaviate")
|
||||
}
|
||||
case types.SQLiteRetrieverEngineType:
|
||||
// No connection config needed for SQLite
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,228 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"github.com/Tencent/WeKnora/internal/errors"
|
||||
"github.com/Tencent/WeKnora/internal/logger"
|
||||
"github.com/Tencent/WeKnora/internal/types"
|
||||
_ "github.com/jackc/pgx/v5/stdlib" // pgx driver for database/sql
|
||||
"github.com/qdrant/go-client/qdrant"
|
||||
"github.com/weaviate/weaviate-go-client/v5/weaviate"
|
||||
"github.com/weaviate/weaviate-go-client/v5/weaviate/auth"
|
||||
wgrpc "github.com/weaviate/weaviate-go-client/v5/weaviate/grpc"
|
||||
)
|
||||
|
||||
const connectionTestTimeout = 10 * time.Second
|
||||
|
||||
// TestConnection tests connectivity to a vector database.
|
||||
// Returns the detected server version on success (e.g., "7.10.1"), empty string if unknown.
|
||||
func (s *vectorStoreService) TestConnection(
|
||||
ctx context.Context,
|
||||
engineType types.RetrieverEngineType,
|
||||
config types.ConnectionConfig,
|
||||
) (string, error) {
|
||||
switch engineType {
|
||||
case types.ElasticsearchRetrieverEngineType:
|
||||
return testElasticsearchConnection(ctx, config)
|
||||
case types.PostgresRetrieverEngineType:
|
||||
return testPostgresConnection(ctx, config)
|
||||
case types.QdrantRetrieverEngineType:
|
||||
return testQdrantConnection(ctx, config)
|
||||
case types.MilvusRetrieverEngineType:
|
||||
return testMilvusConnection(ctx, config)
|
||||
case types.WeaviateRetrieverEngineType:
|
||||
return testWeaviateConnection(ctx, config)
|
||||
case types.SQLiteRetrieverEngineType:
|
||||
// SQLite is file-based, no remote connection to test
|
||||
return "", nil
|
||||
default:
|
||||
return "", errors.NewBadRequestError(
|
||||
fmt.Sprintf("connection test not supported for engine type: %s", engineType))
|
||||
}
|
||||
}
|
||||
|
||||
func testElasticsearchConnection(ctx context.Context, config types.ConnectionConfig) (string, error) {
|
||||
// Use plain HTTP GET to the root endpoint instead of the go-elasticsearch SDK.
|
||||
// The v8 SDK's TypedClient performs a product check that rejects ES7 servers,
|
||||
// so we use a raw HTTP request to support both v7 and v8.
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, config.Addr, nil)
|
||||
if err != nil {
|
||||
return "", errors.NewBadRequestError("failed to create elasticsearch request: invalid address")
|
||||
}
|
||||
if config.Username != "" {
|
||||
req.SetBasicAuth(config.Username, config.Password)
|
||||
}
|
||||
|
||||
client := &http.Client{Timeout: connectionTestTimeout}
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
logger.Warnf(ctx, "Elasticsearch connection test failed: %v", err)
|
||||
return "", errors.NewBadRequestError("failed to connect to elasticsearch: connection refused or authentication failed")
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
logger.Warnf(ctx, "Elasticsearch connection test returned status %d", resp.StatusCode)
|
||||
return "", errors.NewBadRequestError("failed to connect to elasticsearch: authentication failed or server error")
|
||||
}
|
||||
|
||||
// Parse version from response: {"version": {"number": "7.10.1"}, ...}
|
||||
body, err := io.ReadAll(io.LimitReader(resp.Body, 4096))
|
||||
if err != nil {
|
||||
return "", nil // connected but version unknown
|
||||
}
|
||||
|
||||
var esInfo struct {
|
||||
Version struct {
|
||||
Number string `json:"number"`
|
||||
} `json:"version"`
|
||||
}
|
||||
if err := json.Unmarshal(body, &esInfo); err != nil {
|
||||
return "", nil // connected but version unparseable
|
||||
}
|
||||
|
||||
return esInfo.Version.Number, nil
|
||||
}
|
||||
|
||||
func testPostgresConnection(ctx context.Context, config types.ConnectionConfig) (string, error) {
|
||||
testCtx, cancel := context.WithTimeout(ctx, connectionTestTimeout)
|
||||
defer cancel()
|
||||
|
||||
if config.UseDefaultConnection {
|
||||
// Using the default app DB connection — always reachable if the app is running.
|
||||
// Cannot query version without a DB handle; return empty.
|
||||
return "", nil
|
||||
}
|
||||
|
||||
db, err := sql.Open("pgx", config.Addr)
|
||||
if err != nil {
|
||||
return "", errors.NewBadRequestError("failed to create postgres connection: invalid configuration")
|
||||
}
|
||||
defer db.Close()
|
||||
|
||||
if err := db.PingContext(testCtx); err != nil {
|
||||
logger.Warnf(ctx, "Postgres connection test failed: %v", err)
|
||||
return "", errors.NewBadRequestError("failed to connect to postgres: connection refused or authentication failed")
|
||||
}
|
||||
|
||||
// Detect version
|
||||
var version string
|
||||
if err := db.QueryRowContext(testCtx, "SHOW server_version").Scan(&version); err != nil {
|
||||
logger.Warnf(ctx, "Postgres version detection failed: %v", err)
|
||||
return "", nil // connected but version unknown
|
||||
}
|
||||
|
||||
return version, nil
|
||||
}
|
||||
|
||||
func testQdrantConnection(ctx context.Context, config types.ConnectionConfig) (string, error) {
|
||||
testCtx, cancel := context.WithTimeout(ctx, connectionTestTimeout)
|
||||
defer cancel()
|
||||
|
||||
port := config.Port
|
||||
if port == 0 {
|
||||
port = 6334
|
||||
}
|
||||
|
||||
client, err := qdrant.NewClient(&qdrant.Config{
|
||||
Host: config.Host,
|
||||
Port: port,
|
||||
APIKey: config.APIKey,
|
||||
UseTLS: config.UseTLS,
|
||||
})
|
||||
if err != nil {
|
||||
return "", errors.NewBadRequestError("failed to create qdrant client: invalid configuration")
|
||||
}
|
||||
defer client.Close()
|
||||
|
||||
result, err := client.HealthCheck(testCtx)
|
||||
if err != nil {
|
||||
logger.Warnf(ctx, "Qdrant connection test failed: %v", err)
|
||||
return "", errors.NewBadRequestError("failed to connect to qdrant: connection refused or authentication failed")
|
||||
}
|
||||
|
||||
return result.GetVersion(), nil
|
||||
}
|
||||
|
||||
func testMilvusConnection(ctx context.Context, config types.ConnectionConfig) (string, error) {
|
||||
// Use TCP dial instead of the Milvus SDK to avoid protobuf namespace conflict
|
||||
// between milvus-proto and qdrant-client (both register "common.proto").
|
||||
// A TCP dial is sufficient for connectivity verification; the Milvus SDK client
|
||||
// creation in container.go (PR 3) will validate full gRPC connectivity.
|
||||
// Version detection is not possible with TCP dial alone.
|
||||
testCtx, cancel := context.WithTimeout(ctx, connectionTestTimeout)
|
||||
defer cancel()
|
||||
|
||||
addr := config.Addr
|
||||
if addr == "" {
|
||||
addr = "localhost:19530"
|
||||
}
|
||||
|
||||
conn, err := (&net.Dialer{}).DialContext(testCtx, "tcp", addr)
|
||||
if err != nil {
|
||||
logger.Warnf(ctx, "Milvus connection test failed: %v", err)
|
||||
return "", errors.NewBadRequestError("failed to connect to milvus: connection refused or server unreachable")
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
return "", nil
|
||||
}
|
||||
|
||||
func testWeaviateConnection(ctx context.Context, config types.ConnectionConfig) (string, error) {
|
||||
testCtx, cancel := context.WithTimeout(ctx, connectionTestTimeout)
|
||||
defer cancel()
|
||||
|
||||
host := config.Host
|
||||
if host == "" {
|
||||
host = "weaviate:8080"
|
||||
}
|
||||
grpcAddress := config.GrpcAddress
|
||||
if grpcAddress == "" {
|
||||
grpcAddress = "weaviate:50051"
|
||||
}
|
||||
scheme := config.Scheme
|
||||
if scheme == "" {
|
||||
scheme = "http"
|
||||
}
|
||||
|
||||
weaviateCfg := weaviate.Config{
|
||||
Host: host,
|
||||
GrpcConfig: &wgrpc.Config{
|
||||
Host: grpcAddress,
|
||||
},
|
||||
Scheme: scheme,
|
||||
}
|
||||
if config.APIKey != "" {
|
||||
weaviateCfg.AuthConfig = auth.ApiKey{Value: config.APIKey}
|
||||
}
|
||||
|
||||
// Weaviate Go client v5 does not expose a Close() method;
|
||||
// it uses HTTP + gRPC transports that are managed internally.
|
||||
client, err := weaviate.NewClient(weaviateCfg)
|
||||
if err != nil {
|
||||
logger.Warnf(ctx, "Weaviate connection test failed: %v", err)
|
||||
return "", errors.NewBadRequestError("failed to create weaviate client: invalid configuration")
|
||||
}
|
||||
|
||||
isReady, err := client.Misc().ReadyChecker().Do(testCtx)
|
||||
if err != nil || !isReady {
|
||||
logger.Warnf(ctx, "Weaviate connection test failed: ready=%v, err=%v", isReady, err)
|
||||
return "", errors.NewBadRequestError("failed to connect to weaviate: server not ready or authentication failed")
|
||||
}
|
||||
|
||||
// Detect version via /v1/meta
|
||||
meta, err := client.Misc().MetaGetter().Do(testCtx)
|
||||
if err != nil || meta == nil {
|
||||
return "", nil // connected but version unknown
|
||||
}
|
||||
|
||||
return meta.Version, nil
|
||||
}
|
||||
@@ -0,0 +1,525 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/Tencent/WeKnora/internal/errors"
|
||||
"github.com/Tencent/WeKnora/internal/types"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Mock repository
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
type mockVectorStoreRepo struct {
|
||||
stores []*types.VectorStore
|
||||
createErr error
|
||||
updateErr error
|
||||
deleteErr error
|
||||
existsByEndpointErr error
|
||||
existsByEndpoint bool
|
||||
}
|
||||
|
||||
func (m *mockVectorStoreRepo) Create(_ context.Context, store *types.VectorStore) error {
|
||||
if m.createErr != nil {
|
||||
return m.createErr
|
||||
}
|
||||
m.stores = append(m.stores, store)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *mockVectorStoreRepo) GetByID(_ context.Context, tenantID uint64, id string) (*types.VectorStore, error) {
|
||||
for _, s := range m.stores {
|
||||
if s.ID == id && s.TenantID == tenantID {
|
||||
return s, nil
|
||||
}
|
||||
}
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (m *mockVectorStoreRepo) List(_ context.Context, tenantID uint64) ([]*types.VectorStore, error) {
|
||||
var result []*types.VectorStore
|
||||
for _, s := range m.stores {
|
||||
if s.TenantID == tenantID {
|
||||
result = append(result, s)
|
||||
}
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (m *mockVectorStoreRepo) Update(_ context.Context, store *types.VectorStore) error {
|
||||
return m.updateErr
|
||||
}
|
||||
|
||||
func (m *mockVectorStoreRepo) UpdateConnectionConfig(_ context.Context, _ *types.VectorStore) error {
|
||||
return m.updateErr
|
||||
}
|
||||
|
||||
func (m *mockVectorStoreRepo) Delete(_ context.Context, _ uint64, _ string) error {
|
||||
return m.deleteErr
|
||||
}
|
||||
|
||||
func (m *mockVectorStoreRepo) ExistsByEndpointAndIndex(
|
||||
_ context.Context, _ uint64, _ types.RetrieverEngineType, _ string, _ string,
|
||||
) (bool, error) {
|
||||
if m.existsByEndpointErr != nil {
|
||||
return false, m.existsByEndpointErr
|
||||
}
|
||||
return m.existsByEndpoint, nil
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// CreateStore tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestCreateStore_Success(t *testing.T) {
|
||||
repo := &mockVectorStoreRepo{}
|
||||
svc := NewVectorStoreService(repo)
|
||||
|
||||
store := &types.VectorStore{
|
||||
TenantID: 1,
|
||||
Name: "test-es",
|
||||
EngineType: types.ElasticsearchRetrieverEngineType,
|
||||
ConnectionConfig: types.ConnectionConfig{
|
||||
Addr: "http://es:9200",
|
||||
},
|
||||
}
|
||||
|
||||
err := svc.CreateStore(context.Background(), store)
|
||||
assert.NoError(t, err)
|
||||
assert.Len(t, repo.stores, 1)
|
||||
}
|
||||
|
||||
func TestCreateStore_ValidationError(t *testing.T) {
|
||||
repo := &mockVectorStoreRepo{}
|
||||
svc := NewVectorStoreService(repo)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
store *types.VectorStore
|
||||
}{
|
||||
{
|
||||
name: "empty name",
|
||||
store: &types.VectorStore{TenantID: 1, EngineType: types.PostgresRetrieverEngineType},
|
||||
},
|
||||
{
|
||||
name: "invalid engine type",
|
||||
store: &types.VectorStore{TenantID: 1, Name: "test", EngineType: "unknown"},
|
||||
},
|
||||
{
|
||||
name: "zero tenant ID",
|
||||
store: &types.VectorStore{Name: "test", EngineType: types.PostgresRetrieverEngineType},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
err := svc.CreateStore(context.Background(), tt.store)
|
||||
require.Error(t, err)
|
||||
var appErr *errors.AppError
|
||||
assert.ErrorAs(t, err, &appErr)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateStore_ConnectionConfigValidation(t *testing.T) {
|
||||
repo := &mockVectorStoreRepo{}
|
||||
svc := NewVectorStoreService(repo)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
store *types.VectorStore
|
||||
wantError bool
|
||||
}{
|
||||
{
|
||||
name: "elasticsearch without addr",
|
||||
store: &types.VectorStore{
|
||||
TenantID: 1, Name: "test",
|
||||
EngineType: types.ElasticsearchRetrieverEngineType,
|
||||
ConnectionConfig: types.ConnectionConfig{},
|
||||
},
|
||||
wantError: true,
|
||||
},
|
||||
{
|
||||
name: "postgres without addr or default connection",
|
||||
store: &types.VectorStore{
|
||||
TenantID: 1, Name: "test",
|
||||
EngineType: types.PostgresRetrieverEngineType,
|
||||
ConnectionConfig: types.ConnectionConfig{},
|
||||
},
|
||||
wantError: true,
|
||||
},
|
||||
{
|
||||
name: "postgres with use_default_connection",
|
||||
store: &types.VectorStore{
|
||||
TenantID: 1, Name: "test",
|
||||
EngineType: types.PostgresRetrieverEngineType,
|
||||
ConnectionConfig: types.ConnectionConfig{UseDefaultConnection: true},
|
||||
},
|
||||
wantError: false,
|
||||
},
|
||||
{
|
||||
name: "qdrant without host",
|
||||
store: &types.VectorStore{
|
||||
TenantID: 1, Name: "test",
|
||||
EngineType: types.QdrantRetrieverEngineType,
|
||||
ConnectionConfig: types.ConnectionConfig{},
|
||||
},
|
||||
wantError: true,
|
||||
},
|
||||
{
|
||||
name: "milvus without addr",
|
||||
store: &types.VectorStore{
|
||||
TenantID: 1, Name: "test",
|
||||
EngineType: types.MilvusRetrieverEngineType,
|
||||
ConnectionConfig: types.ConnectionConfig{},
|
||||
},
|
||||
wantError: true,
|
||||
},
|
||||
{
|
||||
name: "weaviate without host",
|
||||
store: &types.VectorStore{
|
||||
TenantID: 1, Name: "test",
|
||||
EngineType: types.WeaviateRetrieverEngineType,
|
||||
ConnectionConfig: types.ConnectionConfig{},
|
||||
},
|
||||
wantError: true,
|
||||
},
|
||||
{
|
||||
name: "sqlite with no config (ok)",
|
||||
store: &types.VectorStore{
|
||||
TenantID: 1, Name: "test",
|
||||
EngineType: types.SQLiteRetrieverEngineType,
|
||||
ConnectionConfig: types.ConnectionConfig{},
|
||||
},
|
||||
wantError: false,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
err := svc.CreateStore(context.Background(), tt.store)
|
||||
if tt.wantError {
|
||||
require.Error(t, err)
|
||||
} else {
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateStore_DuplicateCheck_DBStore(t *testing.T) {
|
||||
repo := &mockVectorStoreRepo{existsByEndpoint: true}
|
||||
svc := NewVectorStoreService(repo)
|
||||
|
||||
store := &types.VectorStore{
|
||||
TenantID: 1,
|
||||
Name: "dup-store",
|
||||
EngineType: types.ElasticsearchRetrieverEngineType,
|
||||
ConnectionConfig: types.ConnectionConfig{
|
||||
Addr: "http://es:9200",
|
||||
},
|
||||
}
|
||||
|
||||
err := svc.CreateStore(context.Background(), store)
|
||||
require.Error(t, err)
|
||||
|
||||
var appErr *errors.AppError
|
||||
require.ErrorAs(t, err, &appErr)
|
||||
assert.Equal(t, errors.ErrConflict, appErr.Code)
|
||||
}
|
||||
|
||||
func TestCreateStore_DuplicateCheck_DBError(t *testing.T) {
|
||||
repo := &mockVectorStoreRepo{
|
||||
existsByEndpointErr: assert.AnError,
|
||||
}
|
||||
svc := NewVectorStoreService(repo)
|
||||
|
||||
store := &types.VectorStore{
|
||||
TenantID: 1,
|
||||
Name: "test",
|
||||
EngineType: types.ElasticsearchRetrieverEngineType,
|
||||
ConnectionConfig: types.ConnectionConfig{
|
||||
Addr: "http://es:9200",
|
||||
},
|
||||
}
|
||||
|
||||
err := svc.CreateStore(context.Background(), store)
|
||||
require.Error(t, err)
|
||||
}
|
||||
|
||||
func TestCreateStore_DuplicateCheck_EnvStore(t *testing.T) {
|
||||
// Set up env to simulate an existing elasticsearch env store
|
||||
t.Setenv("RETRIEVE_DRIVER", "elasticsearch_v8")
|
||||
t.Setenv("ELASTICSEARCH_ADDR", "http://es:9200")
|
||||
t.Setenv("ELASTICSEARCH_USERNAME", "elastic")
|
||||
t.Setenv("ELASTICSEARCH_PASSWORD", "secret")
|
||||
t.Setenv("ELASTICSEARCH_INDEX", "xwrag_default")
|
||||
|
||||
repo := &mockVectorStoreRepo{existsByEndpoint: false} // no DB duplicate
|
||||
svc := NewVectorStoreService(repo)
|
||||
|
||||
store := &types.VectorStore{
|
||||
TenantID: 1,
|
||||
Name: "dup-env-store",
|
||||
EngineType: types.ElasticsearchRetrieverEngineType,
|
||||
ConnectionConfig: types.ConnectionConfig{
|
||||
Addr: "http://es:9200",
|
||||
},
|
||||
IndexConfig: types.IndexConfig{
|
||||
IndexName: "xwrag_default",
|
||||
},
|
||||
}
|
||||
|
||||
err := svc.CreateStore(context.Background(), store)
|
||||
require.Error(t, err)
|
||||
|
||||
var appErr *errors.AppError
|
||||
require.ErrorAs(t, err, &appErr)
|
||||
assert.Equal(t, errors.ErrConflict, appErr.Code)
|
||||
assert.Contains(t, appErr.Error(), "environment variables")
|
||||
}
|
||||
|
||||
func TestCreateStore_DuplicateCheck_EnvStore_DifferentIndex_Allowed(t *testing.T) {
|
||||
// Same endpoint as env store but different index — should be allowed
|
||||
t.Setenv("RETRIEVE_DRIVER", "elasticsearch_v8")
|
||||
t.Setenv("ELASTICSEARCH_ADDR", "http://es:9200")
|
||||
t.Setenv("ELASTICSEARCH_INDEX", "xwrag_default")
|
||||
|
||||
repo := &mockVectorStoreRepo{existsByEndpoint: false}
|
||||
svc := NewVectorStoreService(repo)
|
||||
|
||||
store := &types.VectorStore{
|
||||
TenantID: 1,
|
||||
Name: "different-index",
|
||||
EngineType: types.ElasticsearchRetrieverEngineType,
|
||||
ConnectionConfig: types.ConnectionConfig{
|
||||
Addr: "http://es:9200",
|
||||
},
|
||||
IndexConfig: types.IndexConfig{
|
||||
IndexName: "different_index",
|
||||
},
|
||||
}
|
||||
|
||||
err := svc.CreateStore(context.Background(), store)
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
|
||||
func TestCreateStore_DifferentEndpointSameIndex_Allowed(t *testing.T) {
|
||||
repo := &mockVectorStoreRepo{existsByEndpoint: false}
|
||||
svc := NewVectorStoreService(repo)
|
||||
|
||||
store := &types.VectorStore{
|
||||
TenantID: 1,
|
||||
Name: "new-store",
|
||||
EngineType: types.ElasticsearchRetrieverEngineType,
|
||||
ConnectionConfig: types.ConnectionConfig{
|
||||
Addr: "http://es-new:9200",
|
||||
},
|
||||
IndexConfig: types.IndexConfig{
|
||||
IndexName: "shared_index",
|
||||
},
|
||||
}
|
||||
|
||||
err := svc.CreateStore(context.Background(), store)
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// UpdateStore tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestUpdateStore_Success(t *testing.T) {
|
||||
repo := &mockVectorStoreRepo{}
|
||||
svc := NewVectorStoreService(repo)
|
||||
|
||||
store := &types.VectorStore{
|
||||
ID: "test-id",
|
||||
TenantID: 1,
|
||||
Name: "updated-name",
|
||||
}
|
||||
|
||||
err := svc.UpdateStore(context.Background(), store)
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
|
||||
func TestUpdateStore_ValidationError(t *testing.T) {
|
||||
repo := &mockVectorStoreRepo{}
|
||||
svc := NewVectorStoreService(repo)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
store *types.VectorStore
|
||||
}{
|
||||
{
|
||||
name: "empty name",
|
||||
store: &types.VectorStore{ID: "id", TenantID: 1, Name: ""},
|
||||
},
|
||||
{
|
||||
name: "zero tenant ID",
|
||||
store: &types.VectorStore{ID: "id", TenantID: 0, Name: "test"},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
err := svc.UpdateStore(context.Background(), tt.store)
|
||||
require.Error(t, err)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// DeleteStore tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestDeleteStore_Success(t *testing.T) {
|
||||
repo := &mockVectorStoreRepo{}
|
||||
svc := NewVectorStoreService(repo)
|
||||
|
||||
err := svc.DeleteStore(context.Background(), 1, "test-id")
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
|
||||
func TestDeleteStore_RepoError(t *testing.T) {
|
||||
repo := &mockVectorStoreRepo{deleteErr: assert.AnError}
|
||||
svc := NewVectorStoreService(repo)
|
||||
|
||||
err := svc.DeleteStore(context.Background(), 1, "test-id")
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// TestConnection tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestTestConnection_UnsupportedEngineType(t *testing.T) {
|
||||
repo := &mockVectorStoreRepo{}
|
||||
svc := NewVectorStoreService(repo)
|
||||
|
||||
_, err := svc.TestConnection(context.Background(), "unknown_engine", types.ConnectionConfig{})
|
||||
require.Error(t, err)
|
||||
|
||||
var appErr *errors.AppError
|
||||
require.ErrorAs(t, err, &appErr)
|
||||
assert.Equal(t, errors.ErrBadRequest, appErr.Code)
|
||||
}
|
||||
|
||||
func TestTestConnection_SQLiteAlwaysSucceeds(t *testing.T) {
|
||||
repo := &mockVectorStoreRepo{}
|
||||
svc := NewVectorStoreService(repo)
|
||||
|
||||
version, err := svc.TestConnection(context.Background(), types.SQLiteRetrieverEngineType, types.ConnectionConfig{})
|
||||
assert.NoError(t, err)
|
||||
assert.Empty(t, version)
|
||||
}
|
||||
|
||||
func TestTestConnection_PostgresDefaultConnection(t *testing.T) {
|
||||
repo := &mockVectorStoreRepo{}
|
||||
svc := NewVectorStoreService(repo)
|
||||
|
||||
version, err := svc.TestConnection(context.Background(), types.PostgresRetrieverEngineType,
|
||||
types.ConnectionConfig{UseDefaultConnection: true})
|
||||
assert.NoError(t, err)
|
||||
assert.Empty(t, version) // default connection cannot detect version without DB handle
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// validateConnectionConfig tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestValidateConnectionConfig(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
engineType types.RetrieverEngineType
|
||||
config types.ConnectionConfig
|
||||
wantError bool
|
||||
}{
|
||||
{
|
||||
name: "elasticsearch valid",
|
||||
engineType: types.ElasticsearchRetrieverEngineType,
|
||||
config: types.ConnectionConfig{Addr: "http://es:9200"},
|
||||
wantError: false,
|
||||
},
|
||||
{
|
||||
name: "elasticsearch missing addr",
|
||||
engineType: types.ElasticsearchRetrieverEngineType,
|
||||
config: types.ConnectionConfig{},
|
||||
wantError: true,
|
||||
},
|
||||
{
|
||||
name: "postgres with default connection",
|
||||
engineType: types.PostgresRetrieverEngineType,
|
||||
config: types.ConnectionConfig{UseDefaultConnection: true},
|
||||
wantError: false,
|
||||
},
|
||||
{
|
||||
name: "postgres with addr",
|
||||
engineType: types.PostgresRetrieverEngineType,
|
||||
config: types.ConnectionConfig{Addr: "postgres://host:5432/db"},
|
||||
wantError: false,
|
||||
},
|
||||
{
|
||||
name: "postgres without addr or default",
|
||||
engineType: types.PostgresRetrieverEngineType,
|
||||
config: types.ConnectionConfig{},
|
||||
wantError: true,
|
||||
},
|
||||
{
|
||||
name: "qdrant valid",
|
||||
engineType: types.QdrantRetrieverEngineType,
|
||||
config: types.ConnectionConfig{Host: "qdrant-host"},
|
||||
wantError: false,
|
||||
},
|
||||
{
|
||||
name: "qdrant missing host",
|
||||
engineType: types.QdrantRetrieverEngineType,
|
||||
config: types.ConnectionConfig{},
|
||||
wantError: true,
|
||||
},
|
||||
{
|
||||
name: "milvus valid",
|
||||
engineType: types.MilvusRetrieverEngineType,
|
||||
config: types.ConnectionConfig{Addr: "milvus:19530"},
|
||||
wantError: false,
|
||||
},
|
||||
{
|
||||
name: "milvus missing addr",
|
||||
engineType: types.MilvusRetrieverEngineType,
|
||||
config: types.ConnectionConfig{},
|
||||
wantError: true,
|
||||
},
|
||||
{
|
||||
name: "weaviate valid",
|
||||
engineType: types.WeaviateRetrieverEngineType,
|
||||
config: types.ConnectionConfig{Host: "weaviate:8080"},
|
||||
wantError: false,
|
||||
},
|
||||
{
|
||||
name: "weaviate missing host",
|
||||
engineType: types.WeaviateRetrieverEngineType,
|
||||
config: types.ConnectionConfig{},
|
||||
wantError: true,
|
||||
},
|
||||
{
|
||||
name: "sqlite always valid",
|
||||
engineType: types.SQLiteRetrieverEngineType,
|
||||
config: types.ConnectionConfig{},
|
||||
wantError: false,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
err := validateConnectionConfig(tt.engineType, tt.config)
|
||||
if tt.wantError {
|
||||
assert.Error(t, err)
|
||||
} else {
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -6,6 +6,22 @@ import (
|
||||
"github.com/Tencent/WeKnora/internal/types"
|
||||
)
|
||||
|
||||
// VectorStoreService defines the service interface for vector store management.
|
||||
// Tenant isolation is enforced by the handler layer (getOwnedStore pattern).
|
||||
type VectorStoreService interface {
|
||||
// CreateStore validates and creates a new vector store.
|
||||
CreateStore(ctx context.Context, store *types.VectorStore) error
|
||||
// UpdateStore updates an existing vector store (name only).
|
||||
UpdateStore(ctx context.Context, store *types.VectorStore) error
|
||||
// DeleteStore deletes a vector store by tenant + id.
|
||||
DeleteStore(ctx context.Context, tenantID uint64, id string) error
|
||||
// TestConnection tests connectivity to a vector database.
|
||||
// Returns the detected server version on success (e.g., "7.10.1"), empty string if unknown.
|
||||
TestConnection(ctx context.Context, engineType types.RetrieverEngineType, config types.ConnectionConfig) (string, error)
|
||||
// SaveDetectedVersion updates the connection_config.version for a stored vector store.
|
||||
SaveDetectedVersion(ctx context.Context, store *types.VectorStore, version string) error
|
||||
}
|
||||
|
||||
// VectorStoreRepository defines the repository interface for VectorStore CRUD.
|
||||
type VectorStoreRepository interface {
|
||||
// Create creates a new vector store
|
||||
@@ -16,6 +32,8 @@ type VectorStoreRepository interface {
|
||||
List(ctx context.Context, tenantID uint64) ([]*types.VectorStore, error)
|
||||
// Update updates a vector store (only mutable fields: name)
|
||||
Update(ctx context.Context, store *types.VectorStore) error
|
||||
// UpdateConnectionConfig updates only the connection_config column
|
||||
UpdateConnectionConfig(ctx context.Context, store *types.VectorStore) error
|
||||
// Delete soft-deletes a vector store
|
||||
Delete(ctx context.Context, tenantID uint64, id string) error
|
||||
// ExistsByEndpointAndIndex checks if a store with the same endpoint and index already exists
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"database/sql/driver"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Tencent/WeKnora/internal/errors"
|
||||
@@ -12,6 +13,18 @@ import (
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// EnvStoreIDPrefix is the prefix for virtual env store IDs.
|
||||
const EnvStoreIDPrefix = "__env_"
|
||||
|
||||
// IsEnvStoreID checks if the given ID is an env store virtual ID.
|
||||
func IsEnvStoreID(id string) bool {
|
||||
return strings.HasPrefix(id, EnvStoreIDPrefix)
|
||||
}
|
||||
|
||||
// EnvLookupFunc is a function type for looking up environment variables.
|
||||
// In production: os.Getenv, in tests: custom lookup function.
|
||||
type EnvLookupFunc func(string) string
|
||||
|
||||
// VectorStore represents a configured vector database instance for a tenant.
|
||||
// Each tenant can register multiple VectorStore entries (even of the same engine type)
|
||||
// to support multi-store scenarios (e.g., ES-hot + ES-warm clusters).
|
||||
@@ -100,6 +113,9 @@ type ConnectionConfig struct {
|
||||
Scheme string `yaml:"scheme" json:"scheme,omitempty"`
|
||||
// Postgres
|
||||
UseDefaultConnection bool `yaml:"use_default_connection" json:"use_default_connection,omitempty"`
|
||||
// Version is the detected server version (e.g., "7.10.1", "16.2", "1.12.6").
|
||||
// Auto-populated by TestConnection on successful connectivity check.
|
||||
Version string `yaml:"version" json:"version,omitempty"`
|
||||
}
|
||||
|
||||
// Value implements the driver.Valuer interface.
|
||||
@@ -237,3 +253,253 @@ func (c IndexConfig) GetIndexNameOrDefault(engineType RetrieverEngineType) strin
|
||||
return c.IndexName
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// VectorStoreResponse — API response DTO
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
// VectorStoreResponse is the API response DTO for vector store.
|
||||
// Wraps VectorStore with additional metadata (source, readonly).
|
||||
type VectorStoreResponse struct {
|
||||
VectorStore
|
||||
Source string `json:"source"` // "env" or "user"
|
||||
ReadOnly bool `json:"readonly"` // env stores are read-only
|
||||
}
|
||||
|
||||
// NewVectorStoreResponse creates a response DTO from a VectorStore
|
||||
// with sensitive fields masked.
|
||||
func NewVectorStoreResponse(store *VectorStore, source string, readonly bool) VectorStoreResponse {
|
||||
masked := *store
|
||||
masked.ConnectionConfig = store.ConnectionConfig.MaskSensitiveFields()
|
||||
return VectorStoreResponse{
|
||||
VectorStore: masked,
|
||||
Source: source,
|
||||
ReadOnly: readonly,
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// VectorStore type metadata — for /types endpoint
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
// VectorStoreTypeInfo describes a supported engine type and its configuration schema.
|
||||
type VectorStoreTypeInfo struct {
|
||||
Type string `json:"type"`
|
||||
DisplayName string `json:"display_name"`
|
||||
ConnectionFields []VectorStoreFieldInfo `json:"connection_fields"`
|
||||
IndexFields []VectorStoreFieldInfo `json:"index_fields,omitempty"`
|
||||
}
|
||||
|
||||
// VectorStoreFieldInfo describes a single configuration field.
|
||||
type VectorStoreFieldInfo struct {
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"` // "string", "number", "boolean"
|
||||
Required bool `json:"required"`
|
||||
Sensitive bool `json:"sensitive,omitempty"`
|
||||
Default any `json:"default,omitempty"`
|
||||
Description string `json:"description,omitempty"`
|
||||
}
|
||||
|
||||
// GetVectorStoreTypes returns metadata for all supported engine types.
|
||||
func GetVectorStoreTypes() []VectorStoreTypeInfo {
|
||||
return []VectorStoreTypeInfo{
|
||||
{
|
||||
Type: "elasticsearch",
|
||||
DisplayName: "Elasticsearch (Keywords + Vector)",
|
||||
ConnectionFields: []VectorStoreFieldInfo{
|
||||
{Name: "addr", Type: "string", Required: true, Description: "Elasticsearch URL (e.g., http://localhost:9200)"},
|
||||
{Name: "username", Type: "string", Required: false},
|
||||
{Name: "password", Type: "string", Required: false, Sensitive: true},
|
||||
},
|
||||
IndexFields: []VectorStoreFieldInfo{
|
||||
{Name: "index_name", Type: "string", Required: false, Default: "xwrag_default"},
|
||||
{Name: "number_of_shards", Type: "number", Required: false},
|
||||
{Name: "number_of_replicas", Type: "number", Required: false},
|
||||
},
|
||||
},
|
||||
{
|
||||
Type: "postgres",
|
||||
DisplayName: "PostgreSQL (Keywords + Vector)",
|
||||
ConnectionFields: []VectorStoreFieldInfo{
|
||||
{Name: "use_default_connection", Type: "boolean", Required: false, Default: true,
|
||||
Description: "Use the application's default database connection"},
|
||||
{Name: "addr", Type: "string", Required: false,
|
||||
Description: "PostgreSQL connection string (required if use_default_connection is false)"},
|
||||
{Name: "username", Type: "string", Required: false},
|
||||
{Name: "password", Type: "string", Required: false, Sensitive: true},
|
||||
},
|
||||
},
|
||||
{
|
||||
Type: "qdrant",
|
||||
DisplayName: "Qdrant (Keywords + Vector)",
|
||||
ConnectionFields: []VectorStoreFieldInfo{
|
||||
{Name: "host", Type: "string", Required: true, Description: "Qdrant host"},
|
||||
{Name: "port", Type: "number", Required: false, Default: 6334},
|
||||
{Name: "api_key", Type: "string", Required: false, Sensitive: true},
|
||||
{Name: "use_tls", Type: "boolean", Required: false, Default: false},
|
||||
},
|
||||
IndexFields: []VectorStoreFieldInfo{
|
||||
{Name: "collection_prefix", Type: "string", Required: false, Default: "weknora_embeddings"},
|
||||
},
|
||||
},
|
||||
{
|
||||
Type: "milvus",
|
||||
DisplayName: "Milvus (Keywords + Vector)",
|
||||
ConnectionFields: []VectorStoreFieldInfo{
|
||||
{Name: "addr", Type: "string", Required: true, Description: "Milvus address (e.g., localhost:19530)"},
|
||||
{Name: "username", Type: "string", Required: false},
|
||||
{Name: "password", Type: "string", Required: false, Sensitive: true},
|
||||
},
|
||||
IndexFields: []VectorStoreFieldInfo{
|
||||
{Name: "collection_name", Type: "string", Required: false, Default: "weknora_embeddings"},
|
||||
},
|
||||
},
|
||||
{
|
||||
Type: "weaviate",
|
||||
DisplayName: "Weaviate (Keywords + Vector)",
|
||||
ConnectionFields: []VectorStoreFieldInfo{
|
||||
{Name: "host", Type: "string", Required: true, Description: "Weaviate host (e.g., weaviate:8080)"},
|
||||
{Name: "grpc_address", Type: "string", Required: false, Default: "weaviate:50051"},
|
||||
{Name: "scheme", Type: "string", Required: false, Default: "http"},
|
||||
{Name: "api_key", Type: "string", Required: false, Sensitive: true},
|
||||
},
|
||||
IndexFields: []VectorStoreFieldInfo{
|
||||
{Name: "collection_prefix", Type: "string", Required: false, Default: "WeKnora"},
|
||||
},
|
||||
},
|
||||
{
|
||||
Type: "sqlite",
|
||||
DisplayName: "SQLite (Keywords + Vector)",
|
||||
ConnectionFields: []VectorStoreFieldInfo{},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// BuildEnvVectorStores — virtual stores from RETRIEVE_DRIVER env var
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
// BuildEnvVectorStores builds virtual VectorStore entries from RETRIEVE_DRIVER.
|
||||
// Returns []VectorStore (not VectorStoreResponse) so that business logic (e.g.,
|
||||
// duplicate checking) can use them directly. API responses should wrap them
|
||||
// via NewVectorStoreResponse.
|
||||
//
|
||||
// Pure function — does not call os.Getenv directly.
|
||||
//
|
||||
// Usage:
|
||||
//
|
||||
// types.BuildEnvVectorStores(os.Getenv("RETRIEVE_DRIVER"), os.Getenv)
|
||||
func BuildEnvVectorStores(retrieveDriver string, envLookup EnvLookupFunc) []VectorStore {
|
||||
if retrieveDriver == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
drivers := strings.Split(retrieveDriver, ",")
|
||||
var stores []VectorStore
|
||||
|
||||
for _, driver := range drivers {
|
||||
driver = strings.TrimSpace(driver)
|
||||
if driver == "" {
|
||||
continue
|
||||
}
|
||||
|
||||
store := buildEnvStoreForDriver(driver, envLookup)
|
||||
if store != nil {
|
||||
stores = append(stores, *store)
|
||||
}
|
||||
}
|
||||
return stores
|
||||
}
|
||||
|
||||
// FindEnvVectorStore finds a specific env store by its virtual ID.
|
||||
func FindEnvVectorStore(retrieveDriver string, envLookup EnvLookupFunc, id string) *VectorStore {
|
||||
for _, s := range BuildEnvVectorStores(retrieveDriver, envLookup) {
|
||||
if s.ID == id {
|
||||
return &s
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func buildEnvStoreForDriver(driver string, envLookup EnvLookupFunc) *VectorStore {
|
||||
switch driver {
|
||||
case "postgres":
|
||||
return &VectorStore{
|
||||
ID: "__env_postgres__",
|
||||
Name: "postgres (env)",
|
||||
EngineType: PostgresRetrieverEngineType,
|
||||
ConnectionConfig: ConnectionConfig{
|
||||
UseDefaultConnection: true,
|
||||
},
|
||||
}
|
||||
case "sqlite":
|
||||
return &VectorStore{
|
||||
ID: "__env_sqlite__",
|
||||
Name: "sqlite (env)",
|
||||
EngineType: SQLiteRetrieverEngineType,
|
||||
}
|
||||
case "elasticsearch_v8":
|
||||
return &VectorStore{
|
||||
ID: "__env_elasticsearch_v8__",
|
||||
Name: "elasticsearch v8 (env)",
|
||||
EngineType: ElasticsearchRetrieverEngineType,
|
||||
ConnectionConfig: ConnectionConfig{
|
||||
Addr: envLookup("ELASTICSEARCH_ADDR"),
|
||||
Username: envLookup("ELASTICSEARCH_USERNAME"),
|
||||
Password: envLookup("ELASTICSEARCH_PASSWORD"),
|
||||
},
|
||||
IndexConfig: IndexConfig{
|
||||
IndexName: envLookup("ELASTICSEARCH_INDEX"),
|
||||
},
|
||||
}
|
||||
case "elasticsearch_v7":
|
||||
return &VectorStore{
|
||||
ID: "__env_elasticsearch_v7__",
|
||||
Name: "elasticsearch v7 (env)",
|
||||
EngineType: ElasticsearchRetrieverEngineType,
|
||||
ConnectionConfig: ConnectionConfig{
|
||||
Addr: envLookup("ELASTICSEARCH_ADDR"),
|
||||
Username: envLookup("ELASTICSEARCH_USERNAME"),
|
||||
Password: envLookup("ELASTICSEARCH_PASSWORD"),
|
||||
},
|
||||
IndexConfig: IndexConfig{
|
||||
IndexName: envLookup("ELASTICSEARCH_INDEX"),
|
||||
},
|
||||
}
|
||||
case "qdrant":
|
||||
return &VectorStore{
|
||||
ID: "__env_qdrant__",
|
||||
Name: "qdrant (env)",
|
||||
EngineType: QdrantRetrieverEngineType,
|
||||
ConnectionConfig: ConnectionConfig{
|
||||
Host: envLookup("QDRANT_HOST"),
|
||||
APIKey: envLookup("QDRANT_API_KEY"),
|
||||
},
|
||||
}
|
||||
case "milvus":
|
||||
return &VectorStore{
|
||||
ID: "__env_milvus__",
|
||||
Name: "milvus (env)",
|
||||
EngineType: MilvusRetrieverEngineType,
|
||||
ConnectionConfig: ConnectionConfig{
|
||||
Addr: envLookup("MILVUS_ADDRESS"),
|
||||
Username: envLookup("MILVUS_USERNAME"),
|
||||
Password: envLookup("MILVUS_PASSWORD"),
|
||||
},
|
||||
}
|
||||
case "weaviate":
|
||||
return &VectorStore{
|
||||
ID: "__env_weaviate__",
|
||||
Name: "weaviate (env)",
|
||||
EngineType: WeaviateRetrieverEngineType,
|
||||
ConnectionConfig: ConnectionConfig{
|
||||
Host: envLookup("WEAVIATE_HOST"),
|
||||
GrpcAddress: envLookup("WEAVIATE_GRPC_ADDRESS"),
|
||||
Scheme: envLookup("WEAVIATE_SCHEME"),
|
||||
APIKey: envLookup("WEAVIATE_API_KEY"),
|
||||
},
|
||||
}
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
@@ -10,6 +10,245 @@ import (
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// PR2 additions: env store builder, response DTO, types metadata
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
// mockEnvLookup creates a simple env lookup function from a map.
|
||||
func mockEnvLookup(env map[string]string) EnvLookupFunc {
|
||||
return func(key string) string {
|
||||
return env[key]
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsEnvStoreID(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
id string
|
||||
expected bool
|
||||
}{
|
||||
{"env postgres ID", "__env_postgres__", true},
|
||||
{"env elasticsearch ID", "__env_elasticsearch_v8__", true},
|
||||
{"env prefix only", "__env_", true},
|
||||
{"UUID ID", "550e8400-e29b-41d4-a716-446655440000", false},
|
||||
{"empty string", "", false},
|
||||
{"similar but not prefix", "_env_postgres__", false},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
assert.Equal(t, tt.expected, IsEnvStoreID(tt.id))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildEnvVectorStores(t *testing.T) {
|
||||
envMap := map[string]string{
|
||||
"ELASTICSEARCH_ADDR": "http://es:9200",
|
||||
"ELASTICSEARCH_USERNAME": "elastic",
|
||||
"ELASTICSEARCH_PASSWORD": "secret",
|
||||
"ELASTICSEARCH_INDEX": "my_index",
|
||||
"QDRANT_HOST": "qdrant-host",
|
||||
"QDRANT_API_KEY": "qd-key",
|
||||
"MILVUS_ADDRESS": "milvus:19530",
|
||||
"WEAVIATE_HOST": "weaviate:8080",
|
||||
}
|
||||
lookup := mockEnvLookup(envMap)
|
||||
|
||||
t.Run("empty RETRIEVE_DRIVER returns nil", func(t *testing.T) {
|
||||
stores := BuildEnvVectorStores("", lookup)
|
||||
assert.Nil(t, stores)
|
||||
})
|
||||
|
||||
t.Run("single driver postgres", func(t *testing.T) {
|
||||
stores := BuildEnvVectorStores("postgres", lookup)
|
||||
require.Len(t, stores, 1)
|
||||
assert.Equal(t, "__env_postgres__", stores[0].ID)
|
||||
assert.Equal(t, "postgres (env)", stores[0].Name)
|
||||
assert.Equal(t, PostgresRetrieverEngineType, stores[0].EngineType)
|
||||
assert.True(t, stores[0].ConnectionConfig.UseDefaultConnection)
|
||||
})
|
||||
|
||||
t.Run("multiple drivers", func(t *testing.T) {
|
||||
stores := BuildEnvVectorStores("postgres,elasticsearch_v8", lookup)
|
||||
require.Len(t, stores, 2)
|
||||
assert.Equal(t, "__env_postgres__", stores[0].ID)
|
||||
assert.Equal(t, "__env_elasticsearch_v8__", stores[1].ID)
|
||||
assert.Equal(t, "http://es:9200", stores[1].ConnectionConfig.Addr)
|
||||
assert.Equal(t, "elastic", stores[1].ConnectionConfig.Username)
|
||||
assert.Equal(t, "secret", stores[1].ConnectionConfig.Password) // unmasked
|
||||
assert.Equal(t, "my_index", stores[1].IndexConfig.IndexName)
|
||||
})
|
||||
|
||||
t.Run("env store retains raw password (not masked)", func(t *testing.T) {
|
||||
stores := BuildEnvVectorStores("elasticsearch_v8", lookup)
|
||||
require.Len(t, stores, 1)
|
||||
assert.Equal(t, "secret", stores[0].ConnectionConfig.Password)
|
||||
})
|
||||
|
||||
t.Run("unknown driver is skipped", func(t *testing.T) {
|
||||
stores := BuildEnvVectorStores("postgres,unknown_db", lookup)
|
||||
require.Len(t, stores, 1)
|
||||
assert.Equal(t, "__env_postgres__", stores[0].ID)
|
||||
})
|
||||
|
||||
t.Run("whitespace trimmed", func(t *testing.T) {
|
||||
stores := BuildEnvVectorStores(" postgres , elasticsearch_v8 ", lookup)
|
||||
require.Len(t, stores, 2)
|
||||
})
|
||||
|
||||
t.Run("all supported drivers", func(t *testing.T) {
|
||||
stores := BuildEnvVectorStores("postgres,sqlite,elasticsearch_v8,elasticsearch_v7,qdrant,milvus,weaviate", lookup)
|
||||
require.Len(t, stores, 7)
|
||||
|
||||
ids := make([]string, len(stores))
|
||||
for i, s := range stores {
|
||||
ids[i] = s.ID
|
||||
}
|
||||
assert.Contains(t, ids, "__env_postgres__")
|
||||
assert.Contains(t, ids, "__env_sqlite__")
|
||||
assert.Contains(t, ids, "__env_elasticsearch_v8__")
|
||||
assert.Contains(t, ids, "__env_elasticsearch_v7__")
|
||||
assert.Contains(t, ids, "__env_qdrant__")
|
||||
assert.Contains(t, ids, "__env_milvus__")
|
||||
assert.Contains(t, ids, "__env_weaviate__")
|
||||
})
|
||||
|
||||
t.Run("qdrant env store", func(t *testing.T) {
|
||||
stores := BuildEnvVectorStores("qdrant", lookup)
|
||||
require.Len(t, stores, 1)
|
||||
assert.Equal(t, "qdrant-host", stores[0].ConnectionConfig.Host)
|
||||
assert.Equal(t, "qd-key", stores[0].ConnectionConfig.APIKey)
|
||||
})
|
||||
|
||||
t.Run("milvus env store", func(t *testing.T) {
|
||||
stores := BuildEnvVectorStores("milvus", lookup)
|
||||
require.Len(t, stores, 1)
|
||||
assert.Equal(t, "milvus:19530", stores[0].ConnectionConfig.Addr)
|
||||
})
|
||||
|
||||
t.Run("weaviate env store", func(t *testing.T) {
|
||||
stores := BuildEnvVectorStores("weaviate", lookup)
|
||||
require.Len(t, stores, 1)
|
||||
assert.Equal(t, "weaviate:8080", stores[0].ConnectionConfig.Host)
|
||||
})
|
||||
}
|
||||
|
||||
func TestFindEnvVectorStore(t *testing.T) {
|
||||
lookup := mockEnvLookup(map[string]string{})
|
||||
|
||||
t.Run("found", func(t *testing.T) {
|
||||
store := FindEnvVectorStore("postgres", lookup, "__env_postgres__")
|
||||
require.NotNil(t, store)
|
||||
assert.Equal(t, "__env_postgres__", store.ID)
|
||||
})
|
||||
|
||||
t.Run("not found", func(t *testing.T) {
|
||||
store := FindEnvVectorStore("postgres", lookup, "__env_unknown__")
|
||||
assert.Nil(t, store)
|
||||
})
|
||||
|
||||
t.Run("empty driver returns nil", func(t *testing.T) {
|
||||
store := FindEnvVectorStore("", lookup, "__env_postgres__")
|
||||
assert.Nil(t, store)
|
||||
})
|
||||
}
|
||||
|
||||
func TestNewVectorStoreResponse(t *testing.T) {
|
||||
store := &VectorStore{
|
||||
ID: "test-id",
|
||||
Name: "test-store",
|
||||
EngineType: ElasticsearchRetrieverEngineType,
|
||||
ConnectionConfig: ConnectionConfig{
|
||||
Addr: "http://es:9200",
|
||||
Password: "secret",
|
||||
APIKey: "my-api-key",
|
||||
},
|
||||
}
|
||||
|
||||
t.Run("masks sensitive fields", func(t *testing.T) {
|
||||
resp := NewVectorStoreResponse(store, "user", false)
|
||||
assert.Equal(t, "***", resp.ConnectionConfig.Password)
|
||||
assert.Equal(t, "***", resp.ConnectionConfig.APIKey)
|
||||
assert.Equal(t, "http://es:9200", resp.ConnectionConfig.Addr) // non-sensitive preserved
|
||||
})
|
||||
|
||||
t.Run("preserves source and readonly", func(t *testing.T) {
|
||||
resp := NewVectorStoreResponse(store, "env", true)
|
||||
assert.Equal(t, "env", resp.Source)
|
||||
assert.True(t, resp.ReadOnly)
|
||||
})
|
||||
|
||||
t.Run("does not mutate original store", func(t *testing.T) {
|
||||
_ = NewVectorStoreResponse(store, "user", false)
|
||||
assert.Equal(t, "secret", store.ConnectionConfig.Password)
|
||||
assert.Equal(t, "my-api-key", store.ConnectionConfig.APIKey)
|
||||
})
|
||||
|
||||
t.Run("empty sensitive fields not masked to ***", func(t *testing.T) {
|
||||
noSecret := &VectorStore{
|
||||
ID: "test-id",
|
||||
ConnectionConfig: ConnectionConfig{Addr: "http://es:9200"},
|
||||
}
|
||||
resp := NewVectorStoreResponse(noSecret, "user", false)
|
||||
assert.Equal(t, "", resp.ConnectionConfig.Password)
|
||||
assert.Equal(t, "", resp.ConnectionConfig.APIKey)
|
||||
})
|
||||
}
|
||||
|
||||
func TestGetVectorStoreTypes(t *testing.T) {
|
||||
types := GetVectorStoreTypes()
|
||||
|
||||
t.Run("returns 6 engine types", func(t *testing.T) {
|
||||
assert.Len(t, types, 6)
|
||||
})
|
||||
|
||||
t.Run("type names match engine constants", func(t *testing.T) {
|
||||
typeNames := make([]string, len(types))
|
||||
for i, typ := range types {
|
||||
typeNames[i] = typ.Type
|
||||
}
|
||||
assert.Contains(t, typeNames, "elasticsearch")
|
||||
assert.Contains(t, typeNames, "postgres")
|
||||
assert.Contains(t, typeNames, "qdrant")
|
||||
assert.Contains(t, typeNames, "milvus")
|
||||
assert.Contains(t, typeNames, "weaviate")
|
||||
assert.Contains(t, typeNames, "sqlite")
|
||||
})
|
||||
|
||||
t.Run("elasticsearch has connection and index fields", func(t *testing.T) {
|
||||
var esType VectorStoreTypeInfo
|
||||
for _, typ := range types {
|
||||
if typ.Type == "elasticsearch" {
|
||||
esType = typ
|
||||
break
|
||||
}
|
||||
}
|
||||
assert.NotEmpty(t, esType.ConnectionFields)
|
||||
assert.NotEmpty(t, esType.IndexFields)
|
||||
|
||||
// Check sensitive field marking
|
||||
var passwordField VectorStoreFieldInfo
|
||||
for _, f := range esType.ConnectionFields {
|
||||
if f.Name == "password" {
|
||||
passwordField = f
|
||||
break
|
||||
}
|
||||
}
|
||||
assert.True(t, passwordField.Sensitive)
|
||||
})
|
||||
|
||||
t.Run("sqlite has no connection fields", func(t *testing.T) {
|
||||
var sqliteType VectorStoreTypeInfo
|
||||
for _, typ := range types {
|
||||
if typ.Type == "sqlite" {
|
||||
sqliteType = typ
|
||||
break
|
||||
}
|
||||
}
|
||||
assert.Empty(t, sqliteType.ConnectionFields)
|
||||
})
|
||||
}
|
||||
|
||||
// testAESKey is a 32-byte key for testing AES-GCM encryption.
|
||||
const testAESKey = "01234567890123456789012345678901"
|
||||
|
||||
|
||||
Reference in New Issue
Block a user