Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
20 changes: 20 additions & 0 deletions libsql.go
Original file line number Diff line number Diff line change
Expand Up @@ -437,6 +437,26 @@ type conn struct {
nativePtr C.libsql_connection_t
}

// LoadExtension loads a SQLite extension on this connection. Call it through
// sql.Conn.Raw for each connection in a database/sql pool.
func (c *conn) LoadExtension(path, entryPoint string) error {
pathCString := C.CString(path)
defer C.free(unsafe.Pointer(pathCString))

var entryPointCString *C.char
if entryPoint != "" {
entryPointCString = C.CString(entryPoint)
defer C.free(unsafe.Pointer(entryPointCString))
}

var errMsg *C.char
statusCode := C.libsql_load_extension(c.nativePtr, pathCString, entryPointCString, &errMsg)
if statusCode != 0 {
return libsqlError("failed to load extension", statusCode, errMsg)
}
return nil
}

func (c *conn) Prepare(query string) (sqldriver.Stmt, error) {
return c.PrepareContext(context.Background(), query)
}
Expand Down
41 changes: 41 additions & 0 deletions libsql_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -1187,6 +1187,47 @@ func TestErrorCanNotConnect(t *testing.T) {
}
}

func TestLoadExtension(t *testing.T) {
db, err := sql.Open("libsql", "file:"+t.TempDir()+"/test.db")
if err != nil {
t.Fatal(err)
}
defer db.Close()

conn, err := db.Conn(context.Background())
if err != nil {
t.Fatal(err)
}
defer conn.Close()

err = conn.Raw(func(driverConn any) error {
loader, ok := driverConn.(interface{ LoadExtension(string, string) error })
if !ok {
t.Fatal("driver connection does not support extension loading")
}
if err := loader.LoadExtension("/nonexistent/libsql-test-extension", ""); err == nil {
t.Fatal("expected a load error for a missing extension")
}
if path := os.Getenv("LIBSQL_TEST_EXTENSION_PATH"); path != "" {
return loader.LoadExtension(path, "")
}
return nil
})
if err != nil {
t.Fatal(err)
}

if os.Getenv("LIBSQL_TEST_EXTENSION_PATH") != "" {
var nanos int64
if err := conn.QueryRowContext(context.Background(), "SELECT time_get_nano(time_now())").Scan(&nanos); err != nil {
t.Fatal(err)
}
if nanos <= 0 {
t.Fatalf("unexpected nanosecond timestamp %d", nanos)
}
}
}

func TestExec(t *testing.T) {
runMemoryAndFileTests(t, func(t *testing.T, db *sql.DB) {
if _, err := db.ExecContext(context.Background(), "CREATE TABLE test (id INTEGER, name TEXT)"); err != nil {
Expand Down