feat: generated and identity columns for mssql, sqlite, drizzle, typeorm, bun and gorm readers
- mssql: write computed columns (AS (expr) PERSISTED) and identity; read computed definition and identity - sqlite: write GENERATED ALWAYS AS (expr) STORED; read generated columns via table_xinfo and parse the expression - drizzle: generatedAlwaysAs / generated(Always|ByDefault)AsIdentity, read and write - typeorm: asExpression/generatedType and identity decorators, read and write; decorator scan is now quote-aware - bun, gorm readers: read generated and identity markers - readmes updated
This commit is contained in:
@@ -40,6 +40,7 @@ options := &readers.ReaderOptions{
|
||||
- Uses pure Go driver (modernc.org/sqlite) - no CGo required
|
||||
- Supports both file path and connection string
|
||||
- Auto-increment detection for INTEGER PRIMARY KEY columns
|
||||
- Generated columns (virtual and stored) are read via `PRAGMA table_xinfo`; the expression is parsed from the `CREATE TABLE` SQL
|
||||
- Foreign keys require `PRAGMA foreign_keys = ON` to be set
|
||||
|
||||
## Example Schema
|
||||
|
||||
@@ -0,0 +1,116 @@
|
||||
package sqlite
|
||||
|
||||
import (
|
||||
"regexp"
|
||||
"strings"
|
||||
)
|
||||
|
||||
var generatedAsRegex = regexp.MustCompile(`(?is)\bAS\s*\(`)
|
||||
|
||||
// parseGeneratedExpression extracts the expression of a generated column from a CREATE
|
||||
// TABLE statement, e.g. `full TEXT GENERATED ALWAYS AS (a || b) STORED` yields `a || b`.
|
||||
// It returns "" when the column or its expression cannot be found.
|
||||
func parseGeneratedExpression(createSQL, columnName string) string {
|
||||
open := strings.Index(createSQL, "(")
|
||||
if open < 0 {
|
||||
return ""
|
||||
}
|
||||
|
||||
for _, def := range splitTopLevel(createSQL[open+1:]) {
|
||||
if !strings.EqualFold(firstIdentifier(def), columnName) {
|
||||
continue
|
||||
}
|
||||
loc := generatedAsRegex.FindStringIndex(def)
|
||||
if loc == nil {
|
||||
return ""
|
||||
}
|
||||
return balancedParens(def[loc[1]-1:])
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// splitTopLevel splits a table body on commas that are outside parentheses and quotes,
|
||||
// stopping at the parenthesis that closes the body.
|
||||
func splitTopLevel(body string) []string {
|
||||
var parts []string
|
||||
depth := 0
|
||||
var quote byte
|
||||
start := 0
|
||||
for i := 0; i < len(body); i++ {
|
||||
ch := body[i]
|
||||
if quote != 0 {
|
||||
if ch == quote {
|
||||
quote = 0
|
||||
}
|
||||
continue
|
||||
}
|
||||
switch ch {
|
||||
case '\'', '"', '`':
|
||||
quote = ch
|
||||
case '[':
|
||||
quote = ']'
|
||||
case '(':
|
||||
depth++
|
||||
case ')':
|
||||
if depth == 0 {
|
||||
return append(parts, body[start:i])
|
||||
}
|
||||
depth--
|
||||
case ',':
|
||||
if depth == 0 {
|
||||
parts = append(parts, body[start:i])
|
||||
start = i + 1
|
||||
}
|
||||
}
|
||||
}
|
||||
return append(parts, body[start:])
|
||||
}
|
||||
|
||||
// firstIdentifier returns the first (possibly quoted) identifier of a column definition.
|
||||
func firstIdentifier(def string) string {
|
||||
def = strings.TrimSpace(def)
|
||||
if def == "" {
|
||||
return ""
|
||||
}
|
||||
switch def[0] {
|
||||
case '"', '\'', '`':
|
||||
if end := strings.IndexByte(def[1:], def[0]); end >= 0 {
|
||||
return def[1 : 1+end]
|
||||
}
|
||||
case '[':
|
||||
if end := strings.IndexByte(def, ']'); end >= 0 {
|
||||
return def[1:end]
|
||||
}
|
||||
}
|
||||
if end := strings.IndexAny(def, " \t\r\n"); end >= 0 {
|
||||
return def[:end]
|
||||
}
|
||||
return def
|
||||
}
|
||||
|
||||
// balancedParens returns the text inside the parenthesis group that starts at s[0].
|
||||
func balancedParens(s string) string {
|
||||
depth := 0
|
||||
var quote byte
|
||||
for i := 0; i < len(s); i++ {
|
||||
ch := s[i]
|
||||
if quote != 0 {
|
||||
if ch == quote {
|
||||
quote = 0
|
||||
}
|
||||
continue
|
||||
}
|
||||
switch ch {
|
||||
case '\'', '"':
|
||||
quote = ch
|
||||
case '(':
|
||||
depth++
|
||||
case ')':
|
||||
depth--
|
||||
if depth == 0 {
|
||||
return strings.TrimSpace(s[1:i])
|
||||
}
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
@@ -0,0 +1,64 @@
|
||||
package sqlite
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/readers"
|
||||
)
|
||||
|
||||
func TestParseGeneratedExpression(t *testing.T) {
|
||||
createSQL := `CREATE TABLE people (
|
||||
id INTEGER PRIMARY KEY,
|
||||
"first" TEXT,
|
||||
last TEXT,
|
||||
full_name TEXT GENERATED ALWAYS AS (coalesce("first", '') || ' ' || last) STORED,
|
||||
initials TEXT AS (substr("first", 1, 1) || substr(last, 1, 1)),
|
||||
plain TEXT NOT NULL
|
||||
)`
|
||||
|
||||
tests := []struct {
|
||||
column string
|
||||
want string
|
||||
}{
|
||||
{"full_name", `coalesce("first", '') || ' ' || last`},
|
||||
{"initials", `substr("first", 1, 1) || substr(last, 1, 1)`},
|
||||
{"plain", ""},
|
||||
{"missing", ""},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.column, func(t *testing.T) {
|
||||
assert.Equal(t, tt.want, parseGeneratedExpression(createSQL, tt.column))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestReader_GeneratedColumns(t *testing.T) {
|
||||
dbPath := filepath.Join(t.TempDir(), "gen.db")
|
||||
db, err := sql.Open("sqlite", dbPath)
|
||||
require.NoError(t, err)
|
||||
_, err = db.Exec(`CREATE TABLE people (
|
||||
id INTEGER PRIMARY KEY,
|
||||
first TEXT,
|
||||
last TEXT,
|
||||
full_name TEXT GENERATED ALWAYS AS (first || ' ' || last) STORED,
|
||||
initials TEXT AS (substr(first, 1, 1) || substr(last, 1, 1))
|
||||
)`)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, db.Close())
|
||||
|
||||
got, err := NewReader(&readers.ReaderOptions{FilePath: dbPath}).ReadDatabase()
|
||||
require.NoError(t, err)
|
||||
|
||||
cols := got.Schemas[0].Tables[0].Columns
|
||||
require.Contains(t, cols, "full_name", "generated columns must be read")
|
||||
assert.True(t, cols["full_name"].Generated)
|
||||
assert.Equal(t, "first || ' ' || last", cols["full_name"].GenerationExpression)
|
||||
assert.True(t, cols["initials"].Generated)
|
||||
assert.Equal(t, "substr(first, 1, 1) || substr(last, 1, 1)", cols["initials"].GenerationExpression)
|
||||
assert.False(t, cols["first"].Generated)
|
||||
}
|
||||
@@ -75,7 +75,8 @@ func (r *Reader) queryViews() ([]*models.View, error) {
|
||||
|
||||
// queryColumns retrieves all columns for a given table or view
|
||||
func (r *Reader) queryColumns(tableName string) (map[string]*models.Column, error) {
|
||||
query := fmt.Sprintf("PRAGMA table_info(%s)", tableName)
|
||||
// table_xinfo, unlike table_info, also lists generated columns (hidden = 2 virtual, 3 stored)
|
||||
query := fmt.Sprintf("PRAGMA table_xinfo(%s)", tableName)
|
||||
|
||||
rows, err := r.db.QueryContext(r.ctx, query)
|
||||
if err != nil {
|
||||
@@ -84,24 +85,38 @@ func (r *Reader) queryColumns(tableName string) (map[string]*models.Column, erro
|
||||
defer rows.Close()
|
||||
|
||||
columns := make(map[string]*models.Column)
|
||||
var tableSQL string
|
||||
tableSQLLoaded := false
|
||||
|
||||
for rows.Next() {
|
||||
var cid int
|
||||
var name, dataType string
|
||||
var notNull, pk int
|
||||
var notNull, pk, hidden int
|
||||
var defaultValue *string
|
||||
|
||||
if err := rows.Scan(&cid, &name, &dataType, ¬Null, &defaultValue, &pk); err != nil {
|
||||
if err := rows.Scan(&cid, &name, &dataType, ¬Null, &defaultValue, &pk, &hidden); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Hidden virtual-table columns (hidden = 1) are not part of the schema
|
||||
if hidden == 1 {
|
||||
continue
|
||||
}
|
||||
|
||||
column := models.InitColumn(name, tableName, "main")
|
||||
column.Type = r.mapDataType(strings.ToUpper(dataType))
|
||||
column.NotNull = (notNull == 1)
|
||||
column.IsPrimaryKey = (pk > 0)
|
||||
column.Sequence = uint(cid + 1)
|
||||
|
||||
if defaultValue != nil {
|
||||
if hidden == 2 || hidden == 3 {
|
||||
column.Generated = true
|
||||
if !tableSQLLoaded {
|
||||
tableSQL = r.tableSQL(tableName)
|
||||
tableSQLLoaded = true
|
||||
}
|
||||
column.GenerationExpression = parseGeneratedExpression(tableSQL, name)
|
||||
} else if defaultValue != nil {
|
||||
column.Default = *defaultValue
|
||||
}
|
||||
|
||||
@@ -116,6 +131,16 @@ func (r *Reader) queryColumns(tableName string) (map[string]*models.Column, erro
|
||||
return columns, rows.Err()
|
||||
}
|
||||
|
||||
// tableSQL returns the CREATE TABLE statement of a table, or "" when it cannot be read.
|
||||
func (r *Reader) tableSQL(tableName string) string {
|
||||
var sql string
|
||||
err := r.db.QueryRowContext(r.ctx, `SELECT sql FROM sqlite_master WHERE type = 'table' AND name = ?`, tableName).Scan(&sql)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return sql
|
||||
}
|
||||
|
||||
// isAutoIncrement checks if a column is autoincrement
|
||||
func (r *Reader) isAutoIncrement(tableName, columnName string) bool {
|
||||
// Check sqlite_sequence table or parse CREATE TABLE statement
|
||||
|
||||
Reference in New Issue
Block a user