This repository has been archived on 2026-04-03. You can view files and clone it. You cannot open issues or pull requests or push a commit.
Files
core/server/sandbox/db_module_test.go
Jeffrey Smith 4fed810dd5
Some checks failed
CI/CD / detect-changes (pull_request) Successful in 4s
CI/CD / test-frontend (pull_request) Has been skipped
CI/CD / test-runners (pull_request) Has been skipped
CI/CD / e2e-smoke (pull_request) Has been skipped
CI/CD / build-and-deploy (pull_request) Has been cancelled
CI/CD / test-sqlite (pull_request) Has been cancelled
CI/CD / test-go-pg (pull_request) Has been cancelled
Feat v0.8.3 vector column type (#70)
Add vector(N) column type to db_tables manifests with three-tier
progressive enhancement: native pgvector on Postgres, JSONB fallback
without pgvector, TEXT fallback on SQLite. New db.query_similar()
Starlark builtin with dual-path dispatch.

- parseVectorDim validates 1..4096 dimensions
- mapColType gains hasPgvector parameter for tier selection
- HNSW index auto-created on pgvector backends
- starlarkToGoValue extended with list→JSON serialization
- cosineDistance helper for Go-side fallback computation
- ExtensionHandler gains SetCapabilities for install-time DDL
- Roadmap updated: v0.8.4 docs refresh + surface sizing fix

13 new tests (5 schema + 8 db module), all passing with -race.

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-04-03 09:40:38 +00:00

937 lines
30 KiB
Go

package sandbox
import (
"context"
"database/sql"
"math"
"strings"
"testing"
"go.starlark.net/starlark"
_ "modernc.org/sqlite"
)
// newTestDB creates an in-memory SQLite database with a test extension table.
func newTestDB(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 the extension table the tests will use.
// physicalTable("logs") for package "test-ext" → ext_test_ext_logs
_, err = db.Exec(`CREATE TABLE ext_test_ext_logs (
id TEXT PRIMARY KEY,
message TEXT,
user_id TEXT,
count INTEGER,
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,
}
return db, cfg
}
// execScript runs a Starlark snippet with the db module injected.
func execScript(t *testing.T, cfg DBModuleConfig, script string) (*Result, error) {
t.Helper()
sb := New(DefaultConfig())
ctx := context.Background()
modules := map[string]starlark.Value{
"db": BuildDBModule(ctx, cfg),
}
return sb.Exec(ctx, "test.star", script, modules)
}
// ── physicalTable ───────────────────────────
func TestPhysicalTable_ValidName(t *testing.T) {
cfg := DBModuleConfig{PackageID: "my-ext"}
got, err := cfg.physicalTable("logs")
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if got != "ext_my_ext_logs" {
t.Errorf("got %q, want ext_my_ext_logs", got)
}
}
func TestPhysicalTable_EmptyName(t *testing.T) {
cfg := DBModuleConfig{PackageID: "my-ext"}
_, err := cfg.physicalTable("")
if err == nil {
t.Fatal("expected error for empty table name")
}
}
func TestPhysicalTable_InvalidChars(t *testing.T) {
cfg := DBModuleConfig{PackageID: "my-ext"}
_, err := cfg.physicalTable("bad name")
if err == nil {
t.Fatal("expected error for table name with space")
}
_, err = cfg.physicalTable("drop;table")
if err == nil {
t.Fatal("expected error for table name with semicolon")
}
}
// ── db.insert / db.query ────────────────────
func TestDBInsertAndQuery(t *testing.T) {
_, cfg := newTestDB(t)
result, err := execScript(t, cfg, `
row = db.insert("logs", {"message": "hello", "user_id": "u1"})
rows = db.query("logs")
`)
if err != nil {
t.Fatalf("script error: %v", err)
}
rowsVal, ok := result.Globals["rows"]
if !ok {
t.Fatal("rows not in globals")
}
if _, ok := rowsVal.(*starlark.List); !ok {
t.Fatalf("rows is %T, want *starlark.List", rowsVal)
}
}
func TestDBQueryFilters(t *testing.T) {
db, cfg := newTestDB(t)
// Pre-populate with two rows
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, user_id) VALUES ('a', 'one', 'u1')`)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, user_id) VALUES ('b', 'two', 'u2')`)
result, err := execScript(t, cfg, `
rows = db.query("logs", filters={"user_id": "u1"})
`)
if err != nil {
t.Fatalf("script error: %v", err)
}
rowsStr := result.Globals["rows"].String()
if !strings.Contains(rowsStr, "one") {
t.Errorf("expected row with message=one, got: %s", rowsStr)
}
if strings.Contains(rowsStr, "two") {
t.Errorf("unexpected row with message=two: %s", rowsStr)
}
}
func TestDBQueryLimit(t *testing.T) {
db, cfg := newTestDB(t)
for i := 0; i < 5; i++ {
db.Exec(`INSERT INTO ext_test_ext_logs (id, message) VALUES (?, ?)`,
generateID(), "msg")
}
result, err := execScript(t, cfg, `rows = db.query("logs", limit=2)`)
if err != nil {
t.Fatalf("script error: %v", err)
}
// String representation of a 2-element list has exactly 2 entries
rowsStr := result.Globals["rows"].String()
// Count occurrences of "id" key — each row has one
count := strings.Count(rowsStr, `"id"`)
if count != 2 {
t.Errorf("expected 2 rows with limit=2, got repr: %s", rowsStr)
}
}
// ── db.update ───────────────────────────────
func TestDBUpdate(t *testing.T) {
db, cfg := newTestDB(t)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message) VALUES ('r1', 'original')`)
_, err := execScript(t, cfg, `db.update("logs", "r1", {"message": "updated"})`)
if err != nil {
t.Fatalf("script error: %v", err)
}
var msg string
db.QueryRow(`SELECT message FROM ext_test_ext_logs WHERE id = 'r1'`).Scan(&msg)
if msg != "updated" {
t.Errorf("expected message=updated, got %q", msg)
}
}
// ── db.delete ───────────────────────────────
func TestDBDelete(t *testing.T) {
db, cfg := newTestDB(t)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message) VALUES ('r1', 'to-delete')`)
_, err := execScript(t, cfg, `db.delete("logs", "r1")`)
if err != nil {
t.Fatalf("script error: %v", err)
}
var count int
db.QueryRow(`SELECT COUNT(*) FROM ext_test_ext_logs WHERE id = 'r1'`).Scan(&count)
if count != 0 {
t.Errorf("expected row deleted, count = %d", count)
}
}
// ── write guard ─────────────────────────────
func TestDBReadOnlyModuleNoWriteFunctions(t *testing.T) {
_, cfg := newTestDB(t)
cfg.CanWrite = false
mod := BuildDBModule(context.Background(), cfg)
modStr := mod.String()
// insert/update/delete should not be present in a read-only module
for _, fn := range []string{"insert", "update", "delete"} {
if strings.Contains(modStr, fn) {
t.Errorf("read-only module exposes write function %q", fn)
}
}
}
func TestDBWriteFunctionBlockedByReadOnly(t *testing.T) {
_, cfg := newTestDB(t)
cfg.CanWrite = false
_, err := execScript(t, cfg, `db.insert("logs", {"message": "x"})`)
if err == nil {
t.Fatal("expected error calling insert on read-only module")
}
}
// ── view allowlist ───────────────────────────
func TestDBViewRejectsUnknownView(t *testing.T) {
_, cfg := newTestDB(t)
_, err := execScript(t, cfg, `db.view("secrets")`)
if err == nil {
t.Fatal("expected error for unknown view")
}
if !strings.Contains(err.Error(), "unknown view") {
t.Errorf("unexpected error: %v", err)
}
}
// ── before / after range parameters ──────────
func TestDBQueryBefore(t *testing.T) {
db, cfg := newTestDB(t)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, created_at) VALUES ('a', 'old', '2026-01-01T00:00:00Z')`)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, created_at) VALUES ('b', 'mid', '2026-02-01T00:00:00Z')`)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, created_at) VALUES ('c', 'new', '2026-03-01T00:00:00Z')`)
result, err := execScript(t, cfg, `rows = db.query("logs", before={"created_at": "2026-02-15T00:00:00Z"}, order="created_at")`)
if err != nil {
t.Fatalf("script error: %v", err)
}
rowsStr := result.Globals["rows"].String()
if !strings.Contains(rowsStr, "old") || !strings.Contains(rowsStr, "mid") {
t.Errorf("expected old and mid rows, got: %s", rowsStr)
}
if strings.Contains(rowsStr, "new") {
t.Errorf("unexpected 'new' row in before results: %s", rowsStr)
}
}
func TestDBQueryAfter(t *testing.T) {
db, cfg := newTestDB(t)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, created_at) VALUES ('a', 'old', '2026-01-01T00:00:00Z')`)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, created_at) VALUES ('b', 'mid', '2026-02-01T00:00:00Z')`)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, created_at) VALUES ('c', 'new', '2026-03-01T00:00:00Z')`)
result, err := execScript(t, cfg, `rows = db.query("logs", after={"created_at": "2026-02-01T00:00:00Z"}, order="created_at")`)
if err != nil {
t.Fatalf("script error: %v", err)
}
rowsStr := result.Globals["rows"].String()
if !strings.Contains(rowsStr, "new") {
t.Errorf("expected 'new' row, got: %s", rowsStr)
}
if strings.Contains(rowsStr, "old") || strings.Contains(rowsStr, "mid") {
t.Errorf("unexpected old/mid rows in after results: %s", rowsStr)
}
}
func TestDBQueryBeforeWithFilters(t *testing.T) {
db, cfg := newTestDB(t)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, user_id, created_at) VALUES ('a', 'u1-old', 'u1', '2026-01-01T00:00:00Z')`)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, user_id, created_at) VALUES ('b', 'u1-new', 'u1', '2026-03-01T00:00:00Z')`)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, user_id, created_at) VALUES ('c', 'u2-old', 'u2', '2026-01-01T00:00:00Z')`)
result, err := execScript(t, cfg, `rows = db.query("logs", filters={"user_id": "u1"}, before={"created_at": "2026-02-01T00:00:00Z"})`)
if err != nil {
t.Fatalf("script error: %v", err)
}
rowsStr := result.Globals["rows"].String()
if !strings.Contains(rowsStr, "u1-old") {
t.Errorf("expected u1-old, got: %s", rowsStr)
}
if strings.Contains(rowsStr, "u1-new") || strings.Contains(rowsStr, "u2-old") {
t.Errorf("unexpected rows: %s", rowsStr)
}
}
func TestDBQueryBeforeAfterCombined(t *testing.T) {
db, cfg := newTestDB(t)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, created_at) VALUES ('a', 'old', '2026-01-01T00:00:00Z')`)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, created_at) VALUES ('b', 'mid', '2026-02-01T00:00:00Z')`)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, created_at) VALUES ('c', 'new', '2026-03-01T00:00:00Z')`)
result, err := execScript(t, cfg, `rows = db.query("logs", after={"created_at": "2026-01-15T00:00:00Z"}, before={"created_at": "2026-02-15T00:00:00Z"})`)
if err != nil {
t.Fatalf("script error: %v", err)
}
rowsStr := result.Globals["rows"].String()
if !strings.Contains(rowsStr, "mid") {
t.Errorf("expected mid row, got: %s", rowsStr)
}
if strings.Contains(rowsStr, "old") || strings.Contains(rowsStr, "new") {
t.Errorf("unexpected old/new rows: %s", rowsStr)
}
}
func TestDBQueryBeforeEmptyDict(t *testing.T) {
_, cfg := newTestDB(t)
// Empty before dict should be a no-op (same as no before)
_, err := execScript(t, cfg, `rows = db.query("logs", before={})`)
if err != nil {
t.Fatalf("empty before dict should not error: %v", err)
}
}
func TestDBQueryBeforeInvalidColumn(t *testing.T) {
_, cfg := newTestDB(t)
_, err := execScript(t, cfg, `rows = db.query("logs", before={"bad name": "x"})`)
if err == nil {
t.Fatal("expected error for invalid column name in before")
}
if !strings.Contains(err.Error(), "invalid column name") {
t.Errorf("unexpected error: %v", err)
}
}
// ── namespace isolation ──────────────────────
func TestDBQueryRejectsOtherExtensionTable(t *testing.T) {
_, cfg := newTestDB(t)
// Attempt to query a table that belongs to a different extension
_, err := execScript(t, cfg, `db.query("../../etc/passwd")`)
if err == nil {
t.Fatal("expected error for path traversal attempt")
}
}
// ── list_tables ─────────────────────────────
func TestDBListTables(t *testing.T) {
_, cfg := newTestDB(t)
result, err := execScript(t, cfg, `names = db.list_tables()`)
if err != nil {
t.Fatalf("script error: %v", err)
}
namesStr := result.Globals["names"].String()
// The test table ext_test_ext_logs should appear as "logs"
if !strings.Contains(namesStr, "logs") {
t.Errorf("expected 'logs' in list_tables output, got: %s", namesStr)
}
}
// ── search_like ─────────────────────────────
func TestDBQuerySearchLike(t *testing.T) {
db, cfg := newTestDB(t)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, user_id) VALUES ('a', 'hello world', 'u1')`)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, user_id) VALUES ('b', 'goodbye world', 'u2')`)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, user_id) VALUES ('c', 'nothing here', 'u3')`)
result, err := execScript(t, cfg, `rows = db.query("logs", search_like={"message": "%hello%"})`)
if err != nil {
t.Fatalf("script error: %v", err)
}
rowsStr := result.Globals["rows"].String()
if !strings.Contains(rowsStr, "hello world") {
t.Errorf("expected hello world row, got: %s", rowsStr)
}
if strings.Contains(rowsStr, "goodbye") || strings.Contains(rowsStr, "nothing") {
t.Errorf("unexpected rows in search_like results: %s", rowsStr)
}
}
func TestDBQuerySearchLikeMultiColumn(t *testing.T) {
db, cfg := newTestDB(t)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, user_id) VALUES ('a', 'alpha', 'u1')`)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, user_id) VALUES ('b', 'beta', 'u-alpha')`)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, user_id) VALUES ('c', 'gamma', 'u3')`)
// Search across message OR user_id — should match both 'a' (message=alpha) and 'b' (user_id=u-alpha)
result, err := execScript(t, cfg, `rows = db.query("logs", search_like={"message": "%alpha%", "user_id": "%alpha%"})`)
if err != nil {
t.Fatalf("script error: %v", err)
}
rowsStr := result.Globals["rows"].String()
if !strings.Contains(rowsStr, "alpha") {
t.Errorf("expected alpha match, got: %s", rowsStr)
}
if strings.Contains(rowsStr, "gamma") {
t.Errorf("unexpected gamma row in multi-column search: %s", rowsStr)
}
}
func TestDBQuerySearchLikeWithFilters(t *testing.T) {
db, cfg := newTestDB(t)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, user_id) VALUES ('a', 'hello from u1', 'u1')`)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, user_id) VALUES ('b', 'hello from u2', 'u2')`)
// Combine search_like with filters — should only match u1's hello
result, err := execScript(t, cfg, `rows = db.query("logs", filters={"user_id": "u1"}, search_like={"message": "%hello%"})`)
if err != nil {
t.Fatalf("script error: %v", err)
}
rowsStr := result.Globals["rows"].String()
if !strings.Contains(rowsStr, "hello from u1") {
t.Errorf("expected u1 hello row, got: %s", rowsStr)
}
if strings.Contains(rowsStr, "hello from u2") {
t.Errorf("unexpected u2 row when filter + search_like combined: %s", rowsStr)
}
}
func TestDBQuerySearchLikeEmptyDict(t *testing.T) {
_, cfg := newTestDB(t)
// Empty search_like dict should be a no-op
_, err := execScript(t, cfg, `rows = db.query("logs", search_like={})`)
if err != nil {
t.Fatalf("empty search_like dict should not error: %v", err)
}
}
func TestDBQuerySearchLikeInvalidColumn(t *testing.T) {
_, cfg := newTestDB(t)
_, err := execScript(t, cfg, `rows = db.query("logs", search_like={"bad name": "%x%"})`)
if err == nil {
t.Fatal("expected error for invalid column name in search_like")
}
if !strings.Contains(err.Error(), "invalid column name") {
t.Errorf("unexpected error: %v", err)
}
}
// ── db.count ────────────────────────────────
func TestDBCount_Basic(t *testing.T) {
db, cfg := newTestDB(t)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, user_id) VALUES ('a', 'one', 'u1')`)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, user_id) VALUES ('b', 'two', 'u2')`)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, user_id) VALUES ('c', 'three', 'u1')`)
result, err := execScript(t, cfg, `n = db.count("logs")`)
if err != nil {
t.Fatalf("script error: %v", err)
}
n, ok := result.Globals["n"].(starlark.Int)
if !ok {
t.Fatalf("n is %T, want starlark.Int", result.Globals["n"])
}
v, _ := n.Int64()
if v != 3 {
t.Errorf("got %d, want 3", v)
}
}
func TestDBCount_WithFilters(t *testing.T) {
db, cfg := newTestDB(t)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, user_id) VALUES ('a', 'one', 'u1')`)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, user_id) VALUES ('b', 'two', 'u2')`)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, user_id) VALUES ('c', 'three', 'u1')`)
result, err := execScript(t, cfg, `n = db.count("logs", filters={"user_id": "u1"})`)
if err != nil {
t.Fatalf("script error: %v", err)
}
n, ok := result.Globals["n"].(starlark.Int)
if !ok {
t.Fatalf("n is %T, want starlark.Int", result.Globals["n"])
}
v, _ := n.Int64()
if v != 2 {
t.Errorf("got %d, want 2", v)
}
}
func TestDBCount_EmptyTable(t *testing.T) {
_, cfg := newTestDB(t)
result, err := execScript(t, cfg, `n = db.count("logs")`)
if err != nil {
t.Fatalf("script error: %v", err)
}
n, ok := result.Globals["n"].(starlark.Int)
if !ok {
t.Fatalf("n is %T, want starlark.Int", result.Globals["n"])
}
v, _ := n.Int64()
if v != 0 {
t.Errorf("got %d, want 0", v)
}
}
// ── db.aggregate ────────────────────────────
func TestDBAggregate_Sum(t *testing.T) {
db, cfg := newTestDB(t)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, user_id, count) VALUES ('a', 'one', 'u1', 10)`)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, user_id, count) VALUES ('b', 'two', 'u2', 20)`)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, user_id, count) VALUES ('c', 'three', 'u1', 30)`)
result, err := execScript(t, cfg, `val = db.aggregate("logs", "count", "sum")`)
if err != nil {
t.Fatalf("script error: %v", err)
}
val, ok := result.Globals["val"].(starlark.Int)
if !ok {
t.Fatalf("val is %T, want starlark.Int", result.Globals["val"])
}
v, _ := val.Int64()
if v != 60 {
t.Errorf("got %d, want 60", v)
}
}
func TestDBAggregate_Avg(t *testing.T) {
db, cfg := newTestDB(t)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, user_id, count) VALUES ('a', 'one', 'u1', 10)`)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, user_id, count) VALUES ('b', 'two', 'u2', 20)`)
result, err := execScript(t, cfg, `val = db.aggregate("logs", "count", "avg")`)
if err != nil {
t.Fatalf("script error: %v", err)
}
// SQLite AVG of two integers (10+20)/2 = 15.0 → returned as int since 15.0 == int64(15)
val, ok := result.Globals["val"].(starlark.Int)
if !ok {
t.Fatalf("val is %T, want starlark.Int", result.Globals["val"])
}
v, _ := val.Int64()
if v != 15 {
t.Errorf("got %d, want 15", v)
}
}
func TestDBAggregate_MinMax(t *testing.T) {
db, cfg := newTestDB(t)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, user_id, count) VALUES ('a', 'one', 'u1', 5)`)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, user_id, count) VALUES ('b', 'two', 'u2', 50)`)
result, err := execScript(t, cfg, `
mn = db.aggregate("logs", "count", "min")
mx = db.aggregate("logs", "count", "max")
`)
if err != nil {
t.Fatalf("script error: %v", err)
}
mn, _ := result.Globals["mn"].(starlark.Int)
mx, _ := result.Globals["mx"].(starlark.Int)
mnV, _ := mn.Int64()
mxV, _ := mx.Int64()
if mnV != 5 {
t.Errorf("min: got %d, want 5", mnV)
}
if mxV != 50 {
t.Errorf("max: got %d, want 50", mxV)
}
}
func TestDBAggregate_WithFilters(t *testing.T) {
db, cfg := newTestDB(t)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, user_id, count) VALUES ('a', 'one', 'u1', 10)`)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, user_id, count) VALUES ('b', 'two', 'u2', 20)`)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, user_id, count) VALUES ('c', 'three', 'u1', 30)`)
result, err := execScript(t, cfg, `val = db.aggregate("logs", "count", "sum", filters={"user_id": "u1"})`)
if err != nil {
t.Fatalf("script error: %v", err)
}
val, ok := result.Globals["val"].(starlark.Int)
if !ok {
t.Fatalf("val is %T, want starlark.Int", result.Globals["val"])
}
v, _ := val.Int64()
if v != 40 {
t.Errorf("got %d, want 40", v)
}
}
func TestDBAggregate_EmptyResult(t *testing.T) {
_, cfg := newTestDB(t)
result, err := execScript(t, cfg, `val = db.aggregate("logs", "count", "sum")`)
if err != nil {
t.Fatalf("script error: %v", err)
}
if result.Globals["val"] != starlark.None {
t.Errorf("expected None for empty table, got %v", result.Globals["val"])
}
}
func TestDBAggregate_InvalidOp(t *testing.T) {
_, cfg := newTestDB(t)
_, err := execScript(t, cfg, `val = db.aggregate("logs", "count", "median")`)
if err == nil {
t.Fatal("expected error for invalid op")
}
if !strings.Contains(err.Error(), "invalid op") {
t.Errorf("unexpected error: %v", err)
}
}
func TestDBAggregate_InvalidColumn(t *testing.T) {
_, cfg := newTestDB(t)
_, err := execScript(t, cfg, `val = db.aggregate("logs", "bad name", "sum")`)
if err == nil {
t.Fatal("expected error for invalid column name")
}
if !strings.Contains(err.Error(), "invalid column name") {
t.Errorf("unexpected error: %v", err)
}
}
// ── db.query_batch ──────────────────────────
func TestDBQueryBatch_Basic(t *testing.T) {
db, cfg := newTestDB(t)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, user_id) VALUES ('a', 'one', 'u1')`)
db.Exec(`INSERT INTO ext_test_ext_logs (id, message, user_id) VALUES ('b', 'two', 'u2')`)
result, err := execScript(t, cfg, `
results = db.query_batch([
{"table": "logs", "filters": {"user_id": "u1"}},
{"table": "logs", "filters": {"user_id": "u2"}},
])
`)
if err != nil {
t.Fatalf("script error: %v", err)
}
results, ok := result.Globals["results"].(*starlark.List)
if !ok {
t.Fatalf("results is %T, want *starlark.List", result.Globals["results"])
}
if results.Len() != 2 {
t.Fatalf("got %d result sets, want 2", results.Len())
}
// Each result set should have 1 row
for i := 0; i < 2; i++ {
rs, ok := results.Index(i).(*starlark.List)
if !ok {
t.Fatalf("result[%d] is %T, want *starlark.List", i, results.Index(i))
}
if rs.Len() != 1 {
t.Errorf("result[%d] has %d rows, want 1", i, rs.Len())
}
}
}
func TestDBQueryBatch_EmptyList(t *testing.T) {
_, cfg := newTestDB(t)
_, err := execScript(t, cfg, `results = db.query_batch([])`)
if err == nil {
t.Fatal("expected error for empty list")
}
if !strings.Contains(err.Error(), "empty") {
t.Errorf("unexpected error: %v", err)
}
}
func TestDBQueryBatch_TooMany(t *testing.T) {
_, cfg := newTestDB(t)
_, err := execScript(t, cfg, `
def run():
specs = []
for i in range(11):
specs.append({"table": "logs"})
return db.query_batch(specs)
results = run()
`)
if err == nil {
t.Fatal("expected error for >10 queries")
}
if !strings.Contains(err.Error(), "max 10") {
t.Errorf("unexpected error: %v", err)
}
}
func TestDBQueryBatch_MissingTable(t *testing.T) {
_, cfg := newTestDB(t)
_, err := execScript(t, cfg, `results = db.query_batch([{"filters": {"user_id": "u1"}}])`)
if err == nil {
t.Fatal("expected error for missing table key")
}
if !strings.Contains(err.Error(), "table") {
t.Errorf("unexpected error: %v", err)
}
}
// ── 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)
}
}