Feat v0.8.3 vector column (#70)
All checks were successful
CI/CD / detect-changes (push) Successful in 4s
CI/CD / test-runners (push) Has been skipped
CI/CD / test-frontend (push) Has been skipped
CI/CD / e2e-smoke (push) Has been skipped
CI/CD / test-sqlite (push) Successful in 2m59s
CI/CD / test-go-pg (push) Successful in 2m58s
CI/CD / build-and-deploy (push) Successful in 1m14s
All checks were successful
CI/CD / detect-changes (push) Successful in 4s
CI/CD / test-runners (push) Has been skipped
CI/CD / test-frontend (push) Has been skipped
CI/CD / e2e-smoke (push) Has been skipped
CI/CD / test-sqlite (push) Successful in 2m59s
CI/CD / test-go-pg (push) Successful in 2m58s
CI/CD / build-and-deploy (push) Successful in 1m14s
Co-authored-by: Jeffrey Smith <jasafpro@gmail.com> Co-committed-by: Jeffrey Smith <jasafpro@gmail.com>
This commit was merged in pull request #70.
This commit is contained in:
@@ -30,7 +30,10 @@ import (
|
||||
"crypto/rand"
|
||||
"database/sql"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"math"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
"go.starlark.net/starlark"
|
||||
@@ -39,10 +42,11 @@ import (
|
||||
|
||||
// DBModuleConfig holds configuration for a db module instance.
|
||||
type DBModuleConfig struct {
|
||||
PackageID string
|
||||
CanWrite bool // true when db.write permission is granted
|
||||
DB *sql.DB
|
||||
IsPostgres bool
|
||||
PackageID string
|
||||
CanWrite bool // true when db.write permission is granted
|
||||
DB *sql.DB
|
||||
IsPostgres bool
|
||||
HasPgvector bool // true when pgvector extension is available on Postgres
|
||||
}
|
||||
|
||||
// allowedViews is the set of platform views extensions may query via db.view().
|
||||
@@ -55,12 +59,13 @@ var allowedViews = map[string]bool{
|
||||
// prefixed with ext_{packageID}_. Write operations are guarded by CanWrite.
|
||||
func BuildDBModule(ctx context.Context, cfg DBModuleConfig) *starlarkstruct.Module {
|
||||
fns := starlark.StringDict{
|
||||
"query": starlark.NewBuiltin("db.query", dbQuery(ctx, cfg)),
|
||||
"count": starlark.NewBuiltin("db.count", dbCount(ctx, cfg)),
|
||||
"aggregate": starlark.NewBuiltin("db.aggregate", dbAggregate(ctx, cfg)),
|
||||
"query_batch": starlark.NewBuiltin("db.query_batch", dbQueryBatch(ctx, cfg)),
|
||||
"view": starlark.NewBuiltin("db.view", dbView(ctx, cfg)),
|
||||
"list_tables": starlark.NewBuiltin("db.list_tables", dbListTables(ctx, cfg)),
|
||||
"query": starlark.NewBuiltin("db.query", dbQuery(ctx, cfg)),
|
||||
"count": starlark.NewBuiltin("db.count", dbCount(ctx, cfg)),
|
||||
"aggregate": starlark.NewBuiltin("db.aggregate", dbAggregate(ctx, cfg)),
|
||||
"query_batch": starlark.NewBuiltin("db.query_batch", dbQueryBatch(ctx, cfg)),
|
||||
"query_similar": starlark.NewBuiltin("db.query_similar", dbQuerySimilar(ctx, cfg)),
|
||||
"view": starlark.NewBuiltin("db.view", dbView(ctx, cfg)),
|
||||
"list_tables": starlark.NewBuiltin("db.list_tables", dbListTables(ctx, cfg)),
|
||||
}
|
||||
|
||||
if cfg.CanWrite {
|
||||
@@ -160,6 +165,20 @@ func starlarkToGoValue(v starlark.Value) (any, error) {
|
||||
return bool(val), nil
|
||||
case starlark.NoneType:
|
||||
return nil, nil
|
||||
case *starlark.List:
|
||||
goSlice := make([]any, val.Len())
|
||||
for i := 0; i < val.Len(); i++ {
|
||||
elem, err := starlarkToGoValue(val.Index(i))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("list[%d]: %w", i, err)
|
||||
}
|
||||
goSlice[i] = elem
|
||||
}
|
||||
b, err := json.Marshal(goSlice)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("list marshal: %w", err)
|
||||
}
|
||||
return string(b), nil
|
||||
default:
|
||||
return nil, fmt.Errorf("unsupported type %s", v.Type())
|
||||
}
|
||||
@@ -926,3 +945,271 @@ func dbDelete(ctx context.Context, cfg DBModuleConfig) func(*starlark.Thread, *s
|
||||
return starlark.True, nil
|
||||
}
|
||||
}
|
||||
|
||||
// ── Vector Similarity ─────────────────────────
|
||||
|
||||
// starlarkListToFloats converts a Starlark list to a Go float64 slice.
|
||||
func starlarkListToFloats(list *starlark.List) ([]float64, error) {
|
||||
n := list.Len()
|
||||
if n == 0 {
|
||||
return nil, fmt.Errorf("vector must not be empty")
|
||||
}
|
||||
out := make([]float64, n)
|
||||
for i := 0; i < n; i++ {
|
||||
switch v := list.Index(i).(type) {
|
||||
case starlark.Float:
|
||||
out[i] = float64(v)
|
||||
case starlark.Int:
|
||||
i64, _ := v.Int64()
|
||||
out[i] = float64(i64)
|
||||
default:
|
||||
return nil, fmt.Errorf("vector[%d]: expected number, got %s", i, list.Index(i).Type())
|
||||
}
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// floatsToVectorString serializes a float64 slice as a JSON array string.
|
||||
// e.g. [0.1, 0.2, 0.3] → "[0.1,0.2,0.3]"
|
||||
func floatsToVectorString(v []float64) string {
|
||||
b, _ := json.Marshal(v)
|
||||
return string(b)
|
||||
}
|
||||
|
||||
// parseVectorJSON parses a JSON array string (or pgvector text) into float64 slice.
|
||||
func parseVectorJSON(s string) ([]float64, error) {
|
||||
var out []float64
|
||||
if err := json.Unmarshal([]byte(s), &out); err != nil {
|
||||
return nil, fmt.Errorf("parse vector: %w", err)
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// cosineDistance computes cosine distance between two vectors.
|
||||
// Returns 0.0 for identical vectors, 1.0 for orthogonal, 2.0 for opposite.
|
||||
func cosineDistance(a, b []float64) float64 {
|
||||
if len(a) != len(b) || len(a) == 0 {
|
||||
return 1.0
|
||||
}
|
||||
var dot, normA, normB float64
|
||||
for i := range a {
|
||||
dot += a[i] * b[i]
|
||||
normA += a[i] * a[i]
|
||||
normB += b[i] * b[i]
|
||||
}
|
||||
if normA == 0 || normB == 0 {
|
||||
return 1.0
|
||||
}
|
||||
return 1.0 - (dot / (math.Sqrt(normA) * math.Sqrt(normB)))
|
||||
}
|
||||
|
||||
// dbQuerySimilar implements db.query_similar(table, column, vector=[], limit=10, filters=None, metric="cosine").
|
||||
// Returns rows ordered by distance with an injected _distance float.
|
||||
//
|
||||
// Three dispatch paths:
|
||||
// - pgvector: native <=> operator with HNSW index
|
||||
// - fallback: fetch rows, compute cosine distance in Go, sort, return top N
|
||||
func dbQuerySimilar(ctx context.Context, cfg DBModuleConfig) func(*starlark.Thread, *starlark.Builtin, starlark.Tuple, []starlark.Tuple) (starlark.Value, error) {
|
||||
return func(thread *starlark.Thread, b *starlark.Builtin, args starlark.Tuple, kwargs []starlark.Tuple) (starlark.Value, error) {
|
||||
var table, column string
|
||||
var vectorVal *starlark.List
|
||||
var limit starlark.Int = starlark.MakeInt(10)
|
||||
var filters starlark.Value = starlark.None
|
||||
var metric starlark.String = "cosine"
|
||||
|
||||
if err := starlark.UnpackArgs(b.Name(), args, kwargs,
|
||||
"table", &table,
|
||||
"column", &column,
|
||||
"vector", &vectorVal,
|
||||
"limit?", &limit,
|
||||
"filters?", &filters,
|
||||
"metric?", &metric,
|
||||
); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if string(metric) != "cosine" {
|
||||
return nil, fmt.Errorf("db.query_similar: unsupported metric %q (only \"cosine\" is supported)", string(metric))
|
||||
}
|
||||
|
||||
queryVec, err := starlarkListToFloats(vectorVal)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("db.query_similar: %w", err)
|
||||
}
|
||||
|
||||
lim, ok := limit.Int64()
|
||||
if !ok || lim < 1 {
|
||||
lim = 10
|
||||
}
|
||||
if lim > 100 {
|
||||
lim = 100
|
||||
}
|
||||
|
||||
if strings.ContainsAny(column, " \t\n\"';-") {
|
||||
return nil, fmt.Errorf("db.query_similar: invalid column name %q", column)
|
||||
}
|
||||
|
||||
physTable, err := cfg.physicalTable(table)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if cfg.HasPgvector {
|
||||
return querySimilarPgvector(ctx, cfg, physTable, column, queryVec, int(lim), filters)
|
||||
}
|
||||
return querySimilarFallback(ctx, cfg, physTable, column, queryVec, int(lim), filters)
|
||||
}
|
||||
}
|
||||
|
||||
// querySimilarPgvector uses the native pgvector <=> operator.
|
||||
func querySimilarPgvector(ctx context.Context, cfg DBModuleConfig, physTable, column string, queryVec []float64, limit int, filters starlark.Value) (starlark.Value, error) {
|
||||
vecStr := floatsToVectorString(queryVec)
|
||||
|
||||
whereClause, whereArgs, err := cfg.starlarkFiltersToSQL(filters, 1)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
nextIdx := len(whereArgs) + 1
|
||||
vecPH := cfg.ph(nextIdx)
|
||||
limitPH := cfg.ph(nextIdx + 1)
|
||||
|
||||
query := fmt.Sprintf("SELECT *, (%s <=> %s::vector) AS _distance FROM %s",
|
||||
column, vecPH, physTable)
|
||||
if whereClause != "" {
|
||||
query += " " + whereClause
|
||||
}
|
||||
query += fmt.Sprintf(" ORDER BY %s <=> %s::vector LIMIT %s", column, vecPH, limitPH)
|
||||
|
||||
allArgs := append(whereArgs, vecStr, limit)
|
||||
|
||||
rows, err := cfg.DB.QueryContext(ctx, query, allArgs...)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("db.query_similar: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
return rowsToStarlark(rows)
|
||||
}
|
||||
|
||||
// querySimilarFallback fetches rows and computes cosine distance in Go.
|
||||
// Used for SQLite (TEXT) and Postgres without pgvector (JSONB).
|
||||
func querySimilarFallback(ctx context.Context, cfg DBModuleConfig, physTable, column string, queryVec []float64, limit int, filters starlark.Value) (starlark.Value, error) {
|
||||
whereClause, whereArgs, err := cfg.starlarkFiltersToSQL(filters, 1)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Fetch up to 1000 rows for distance computation.
|
||||
fetchLimit := limit * 10
|
||||
if fetchLimit > 1000 {
|
||||
fetchLimit = 1000
|
||||
}
|
||||
if fetchLimit < limit {
|
||||
fetchLimit = limit
|
||||
}
|
||||
|
||||
limitPH := cfg.ph(len(whereArgs) + 1)
|
||||
query := fmt.Sprintf("SELECT * FROM %s", physTable)
|
||||
if whereClause != "" {
|
||||
query += " " + whereClause
|
||||
}
|
||||
query += " LIMIT " + limitPH
|
||||
allArgs := append(whereArgs, fetchLimit)
|
||||
|
||||
rows, err := cfg.DB.QueryContext(ctx, query, allArgs...)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("db.query_similar: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
cols, err := rows.Columns()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Find the vector column index.
|
||||
vecColIdx := -1
|
||||
for i, c := range cols {
|
||||
if c == column {
|
||||
vecColIdx = i
|
||||
break
|
||||
}
|
||||
}
|
||||
if vecColIdx == -1 {
|
||||
return nil, fmt.Errorf("db.query_similar: column %q not found in table", column)
|
||||
}
|
||||
|
||||
type scoredRow struct {
|
||||
vals []any
|
||||
distance float64
|
||||
}
|
||||
var scored []scoredRow
|
||||
|
||||
for rows.Next() {
|
||||
vals := make([]any, len(cols))
|
||||
ptrs := make([]any, len(cols))
|
||||
for i := range vals {
|
||||
ptrs[i] = &vals[i]
|
||||
}
|
||||
if err := rows.Scan(ptrs...); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Parse the vector column value.
|
||||
var vecStr string
|
||||
switch v := vals[vecColIdx].(type) {
|
||||
case string:
|
||||
vecStr = v
|
||||
case []byte:
|
||||
vecStr = string(v)
|
||||
default:
|
||||
continue // skip rows with unparseable vectors
|
||||
}
|
||||
|
||||
rowVec, err := parseVectorJSON(vecStr)
|
||||
if err != nil {
|
||||
continue // skip malformed vectors
|
||||
}
|
||||
|
||||
dist := cosineDistance(queryVec, rowVec)
|
||||
scored = append(scored, scoredRow{vals: vals, distance: dist})
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Sort by distance ascending.
|
||||
sort.Slice(scored, func(i, j int) bool {
|
||||
return scored[i].distance < scored[j].distance
|
||||
})
|
||||
|
||||
// Take top N and build Starlark dicts.
|
||||
if len(scored) > limit {
|
||||
scored = scored[:limit]
|
||||
}
|
||||
|
||||
var result []starlark.Value
|
||||
for _, sr := range scored {
|
||||
d := starlark.NewDict(len(cols) + 1)
|
||||
for i, col := range cols {
|
||||
sv, err := goToStarlark(sr.vals[i])
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("column %q: %w", col, err)
|
||||
}
|
||||
if err := d.SetKey(starlark.String(col), sv); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
// Inject _distance.
|
||||
if err := d.SetKey(starlark.String("_distance"), starlark.Float(sr.distance)); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result = append(result, d)
|
||||
}
|
||||
|
||||
if result == nil {
|
||||
result = []starlark.Value{}
|
||||
}
|
||||
return starlark.NewList(result), nil
|
||||
}
|
||||
|
||||
@@ -3,6 +3,7 @@ package sandbox
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"math"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
@@ -684,3 +685,252 @@ func TestDBQueryBatch_MissingTable(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// ── starlarkToGoValue list ──────────────────
|
||||
|
||||
func TestStarlarkToGoValue_List(t *testing.T) {
|
||||
list := starlark.NewList([]starlark.Value{
|
||||
starlark.Float(0.1), starlark.Float(0.2), starlark.Float(0.3),
|
||||
})
|
||||
got, err := starlarkToGoValue(list)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
s, ok := got.(string)
|
||||
if !ok {
|
||||
t.Fatalf("expected string, got %T", got)
|
||||
}
|
||||
if s != "[0.1,0.2,0.3]" {
|
||||
t.Errorf("got %q, want %q", s, "[0.1,0.2,0.3]")
|
||||
}
|
||||
}
|
||||
|
||||
// ── cosineDistance ───────────────────────────
|
||||
|
||||
func TestCosineDistance(t *testing.T) {
|
||||
// Identical vectors → distance 0
|
||||
d := cosineDistance([]float64{1, 0, 0}, []float64{1, 0, 0})
|
||||
if math.Abs(d) > 1e-9 {
|
||||
t.Errorf("identical vectors: got %f, want 0", d)
|
||||
}
|
||||
|
||||
// Orthogonal vectors → distance 1
|
||||
d = cosineDistance([]float64{1, 0}, []float64{0, 1})
|
||||
if math.Abs(d-1.0) > 1e-9 {
|
||||
t.Errorf("orthogonal vectors: got %f, want 1", d)
|
||||
}
|
||||
|
||||
// Opposite vectors → distance 2
|
||||
d = cosineDistance([]float64{1, 0}, []float64{-1, 0})
|
||||
if math.Abs(d-2.0) > 1e-9 {
|
||||
t.Errorf("opposite vectors: got %f, want 2", d)
|
||||
}
|
||||
|
||||
// Empty/mismatched → distance 1
|
||||
d = cosineDistance([]float64{}, []float64{})
|
||||
if d != 1.0 {
|
||||
t.Errorf("empty vectors: got %f, want 1", d)
|
||||
}
|
||||
d = cosineDistance([]float64{1}, []float64{1, 2})
|
||||
if d != 1.0 {
|
||||
t.Errorf("mismatched vectors: got %f, want 1", d)
|
||||
}
|
||||
}
|
||||
|
||||
// ── vector helper: newVectorTestDB ──────────
|
||||
|
||||
func newVectorTestDB(t *testing.T) (*sql.DB, DBModuleConfig) {
|
||||
t.Helper()
|
||||
db, err := sql.Open("sqlite", ":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open db: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { db.Close() })
|
||||
|
||||
// Create a table with a TEXT column for vector storage (SQLite fallback).
|
||||
_, err = db.Exec(`CREATE TABLE ext_test_ext_embeddings (
|
||||
id TEXT PRIMARY KEY,
|
||||
embedding TEXT,
|
||||
category TEXT,
|
||||
created_at TEXT DEFAULT (datetime('now'))
|
||||
)`)
|
||||
if err != nil {
|
||||
t.Fatalf("create test table: %v", err)
|
||||
}
|
||||
|
||||
cfg := DBModuleConfig{
|
||||
PackageID: "test-ext",
|
||||
CanWrite: true,
|
||||
DB: db,
|
||||
IsPostgres: false,
|
||||
HasPgvector: false,
|
||||
}
|
||||
return db, cfg
|
||||
}
|
||||
|
||||
// ── db.query_similar ────────────────────────
|
||||
|
||||
func TestDBQuerySimilar_Basic(t *testing.T) {
|
||||
db, cfg := newVectorTestDB(t)
|
||||
|
||||
// Insert 5 rows with known vectors.
|
||||
db.Exec(`INSERT INTO ext_test_ext_embeddings (id, embedding, category) VALUES ('a', '[1,0,0]', 'x')`)
|
||||
db.Exec(`INSERT INTO ext_test_ext_embeddings (id, embedding, category) VALUES ('b', '[0,1,0]', 'x')`)
|
||||
db.Exec(`INSERT INTO ext_test_ext_embeddings (id, embedding, category) VALUES ('c', '[0,0,1]', 'x')`)
|
||||
db.Exec(`INSERT INTO ext_test_ext_embeddings (id, embedding, category) VALUES ('d', '[0.9,0.1,0]', 'x')`)
|
||||
db.Exec(`INSERT INTO ext_test_ext_embeddings (id, embedding, category) VALUES ('e', '[0,0.9,0.1]', 'x')`)
|
||||
|
||||
// Query similar to [1,0,0] — should return 'a' first (identical), then 'd' (close).
|
||||
result, err := execScript(t, cfg, `
|
||||
rows = db.query_similar("embeddings", "embedding", vector=[1.0, 0.0, 0.0], limit=3)
|
||||
`)
|
||||
if err != nil {
|
||||
t.Fatalf("script error: %v", err)
|
||||
}
|
||||
|
||||
rows, ok := result.Globals["rows"].(*starlark.List)
|
||||
if !ok {
|
||||
t.Fatalf("rows is %T, want *starlark.List", result.Globals["rows"])
|
||||
}
|
||||
if rows.Len() != 3 {
|
||||
t.Fatalf("got %d rows, want 3", rows.Len())
|
||||
}
|
||||
|
||||
// First row should be 'a' (distance ≈ 0).
|
||||
first := rows.Index(0).(*starlark.Dict)
|
||||
idVal, _, _ := first.Get(starlark.String("id"))
|
||||
if string(idVal.(starlark.String)) != "a" {
|
||||
t.Errorf("first row id = %s, want 'a'", idVal)
|
||||
}
|
||||
|
||||
// Verify _distance is present and near 0.
|
||||
distVal, _, _ := first.Get(starlark.String("_distance"))
|
||||
dist := float64(distVal.(starlark.Float))
|
||||
if dist > 0.01 {
|
||||
t.Errorf("first row _distance = %f, want ≈ 0", dist)
|
||||
}
|
||||
|
||||
// Second row should be 'd' (closest after identical).
|
||||
second := rows.Index(1).(*starlark.Dict)
|
||||
idVal2, _, _ := second.Get(starlark.String("id"))
|
||||
if string(idVal2.(starlark.String)) != "d" {
|
||||
t.Errorf("second row id = %s, want 'd'", idVal2)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDBQuerySimilar_WithFilters(t *testing.T) {
|
||||
db, cfg := newVectorTestDB(t)
|
||||
|
||||
db.Exec(`INSERT INTO ext_test_ext_embeddings (id, embedding, category) VALUES ('a', '[1,0,0]', 'alpha')`)
|
||||
db.Exec(`INSERT INTO ext_test_ext_embeddings (id, embedding, category) VALUES ('b', '[0.9,0.1,0]', 'beta')`)
|
||||
db.Exec(`INSERT INTO ext_test_ext_embeddings (id, embedding, category) VALUES ('c', '[0,1,0]', 'alpha')`)
|
||||
|
||||
// Filter to alpha only — should exclude 'b' even though it's close.
|
||||
result, err := execScript(t, cfg, `
|
||||
rows = db.query_similar("embeddings", "embedding", vector=[1.0, 0.0, 0.0], filters={"category": "alpha"})
|
||||
`)
|
||||
if err != nil {
|
||||
t.Fatalf("script error: %v", err)
|
||||
}
|
||||
|
||||
rows := result.Globals["rows"].(*starlark.List)
|
||||
for i := 0; i < rows.Len(); i++ {
|
||||
d := rows.Index(i).(*starlark.Dict)
|
||||
idVal, _, _ := d.Get(starlark.String("id"))
|
||||
if string(idVal.(starlark.String)) == "b" {
|
||||
t.Error("filtered query should not include row 'b' (category=beta)")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestDBQuerySimilar_EmptyTable(t *testing.T) {
|
||||
_, cfg := newVectorTestDB(t)
|
||||
|
||||
result, err := execScript(t, cfg, `
|
||||
rows = db.query_similar("embeddings", "embedding", vector=[1.0, 0.0, 0.0])
|
||||
`)
|
||||
if err != nil {
|
||||
t.Fatalf("script error: %v", err)
|
||||
}
|
||||
rows := result.Globals["rows"].(*starlark.List)
|
||||
if rows.Len() != 0 {
|
||||
t.Errorf("expected 0 rows for empty table, got %d", rows.Len())
|
||||
}
|
||||
}
|
||||
|
||||
func TestDBQuerySimilar_InvalidMetric(t *testing.T) {
|
||||
_, cfg := newVectorTestDB(t)
|
||||
|
||||
_, err := execScript(t, cfg, `
|
||||
rows = db.query_similar("embeddings", "embedding", vector=[1.0], metric="l2")
|
||||
`)
|
||||
if err == nil {
|
||||
t.Fatal("expected error for unsupported metric")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "unsupported metric") {
|
||||
t.Errorf("unexpected error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDBQuerySimilar_LimitCap(t *testing.T) {
|
||||
db, cfg := newVectorTestDB(t)
|
||||
|
||||
// Insert 5 rows
|
||||
for i := 0; i < 5; i++ {
|
||||
db.Exec(`INSERT INTO ext_test_ext_embeddings (id, embedding) VALUES (?, '[1,0,0]')`,
|
||||
generateID())
|
||||
}
|
||||
|
||||
result, err := execScript(t, cfg, `
|
||||
rows = db.query_similar("embeddings", "embedding", vector=[1.0, 0.0, 0.0], limit=2)
|
||||
`)
|
||||
if err != nil {
|
||||
t.Fatalf("script error: %v", err)
|
||||
}
|
||||
rows := result.Globals["rows"].(*starlark.List)
|
||||
if rows.Len() != 2 {
|
||||
t.Errorf("expected 2 rows with limit=2, got %d", rows.Len())
|
||||
}
|
||||
}
|
||||
|
||||
func TestDBInsert_VectorColumn(t *testing.T) {
|
||||
db, cfg := newVectorTestDB(t)
|
||||
|
||||
_, err := execScript(t, cfg, `
|
||||
row = db.insert("embeddings", {"embedding": [0.1, 0.2, 0.3], "category": "test"})
|
||||
`)
|
||||
if err != nil {
|
||||
t.Fatalf("script error: %v", err)
|
||||
}
|
||||
|
||||
// Verify the stored value is valid JSON.
|
||||
var stored string
|
||||
db.QueryRow(`SELECT embedding FROM ext_test_ext_embeddings WHERE category = 'test'`).Scan(&stored)
|
||||
if stored != "[0.1,0.2,0.3]" {
|
||||
t.Errorf("stored vector = %q, want %q", stored, "[0.1,0.2,0.3]")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDBQuerySimilar_InsertThenQuery(t *testing.T) {
|
||||
_, cfg := newVectorTestDB(t)
|
||||
|
||||
// Full round-trip: insert via db.insert, then query_similar.
|
||||
result, err := execScript(t, cfg, `
|
||||
db.insert("embeddings", {"embedding": [1.0, 0.0, 0.0], "category": "a"})
|
||||
db.insert("embeddings", {"embedding": [0.0, 1.0, 0.0], "category": "b"})
|
||||
rows = db.query_similar("embeddings", "embedding", vector=[1.0, 0.0, 0.0], limit=1)
|
||||
`)
|
||||
if err != nil {
|
||||
t.Fatalf("script error: %v", err)
|
||||
}
|
||||
|
||||
rows := result.Globals["rows"].(*starlark.List)
|
||||
if rows.Len() != 1 {
|
||||
t.Fatalf("expected 1 row, got %d", rows.Len())
|
||||
}
|
||||
first := rows.Index(0).(*starlark.Dict)
|
||||
catVal, _, _ := first.Get(starlark.String("category"))
|
||||
if string(catVal.(starlark.String)) != "a" {
|
||||
t.Errorf("closest match should be category=a, got %s", catVal)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -425,10 +425,11 @@ func (r *Runner) buildModulesWithLibCtx(ctx context.Context, packageID string, m
|
||||
// Wire db module at the highest granted level.
|
||||
if dbLevel > 0 && r.db != nil {
|
||||
modules["db"] = BuildDBModule(ctx, DBModuleConfig{
|
||||
PackageID: packageID,
|
||||
CanWrite: dbLevel == 2,
|
||||
DB: r.db,
|
||||
IsPostgres: r.dbPostgres,
|
||||
PackageID: packageID,
|
||||
CanWrite: dbLevel == 2,
|
||||
DB: r.db,
|
||||
IsPostgres: r.dbPostgres,
|
||||
HasPgvector: r.capabilities["pgvector"],
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user