diff --git a/pkg/pgsql/extensions.go b/pkg/pgsql/extensions.go new file mode 100644 index 0000000..ae97a25 --- /dev/null +++ b/pkg/pgsql/extensions.go @@ -0,0 +1,459 @@ +package pgsql + +import ( + "sort" + "strings" +) + +// Extension describes a PostgreSQL extension RelSpec recognizes, along with the schema +// artefacts that imply it: the types it provides (declared on TypeSpec.Extension), the +// index access methods and operator classes it installs, and the functions whose use in a +// default, check constraint, index predicate, or view body requires it. +type Extension struct { + Name string + Category string + Description string + + // Requires lists extensions that must be created before this one. + Requires []string + + // IndexMethods are access methods usable as Index.Type. + IndexMethods []string + + // OperatorClasses are operator classes the extension installs. + OperatorClasses []string + + // Functions are function names whose use implies the extension. + Functions []string + + // FunctionPrefixes match whole families of functions (e.g. "st_" for PostGIS). + FunctionPrefixes []string +} + +// postgresExtensions is the set of extensions RelSpec knows how to detect and emit. +var postgresExtensions = map[string]Extension{ + "amcheck": { + Name: "amcheck", Category: "integrity", + Description: "Verifies B-tree and related structure consistency to help detect corruption.", + Functions: []string{"bt_index_check", "bt_index_parent_check", "verify_heapam"}, + }, + "btree_gin": { + Name: "btree_gin", Category: "indexing", + Description: "Adds GIN operator classes for common scalar data types.", + }, + "btree_gist": { + Name: "btree_gist", Category: "indexing", + Description: "Adds GiST operator classes for common scalar data types and exclusion constraints.", + }, + "citext": { + Name: "citext", Category: "text", + Description: "Provides case-insensitive text columns and operators.", + Functions: []string{"citext"}, + }, + "fuzzystrmatch": { + Name: "fuzzystrmatch", Category: "text", + Description: "Adds phonetic and fuzzy matching helpers like Soundex and Levenshtein.", + Functions: []string{ + "soundex", "difference", "levenshtein", "levenshtein_less_equal", + "metaphone", "dmetaphone", "dmetaphone_alt", + }, + }, + "hstore": { + Name: "hstore", Category: "document", + Description: "Adds a lightweight key/value data type for semi-structured attributes.", + OperatorClasses: []string{"gin_hstore_ops", "gist_hstore_ops", "hash_hstore_ops", "btree_hstore_ops"}, + Functions: []string{ + "hstore", "akeys", "avals", "skeys", "svals", + "hstore_to_json", "hstore_to_jsonb", "hstore_to_array", "hstore_to_matrix", + }, + }, + "http": { + Name: "http", Category: "integration", + Description: "Lets SQL functions make outbound HTTP requests.", + Functions: []string{ + "http", "http_get", "http_post", "http_put", "http_patch", "http_delete", + "http_head", "urlencode", + }, + }, + "pg_background": { + Name: "pg_background", Category: "jobs", + Description: "Runs SQL asynchronously in PostgreSQL background workers.", + Functions: []string{"pg_background_launch", "pg_background_result", "pg_background_detach"}, + }, + "pg_cron": { + Name: "pg_cron", Category: "scheduling", + Description: "Schedules recurring SQL jobs inside PostgreSQL.", + FunctionPrefixes: []string{"cron."}, + }, + "pg_jsonschema": { + Name: "pg_jsonschema", Category: "validation", + Description: "Validates json and jsonb values against JSON Schema.", + Functions: []string{"json_matches_schema", "jsonb_matches_schema", "jsonschema_is_valid"}, + }, + "pg_partman": { + Name: "pg_partman", Category: "partitioning", + Description: "Automates time-based and serial-based partition management.", + FunctionPrefixes: []string{"partman."}, + }, + "pg_qualstats": { + Name: "pg_qualstats", Category: "observability", + Description: "Tracks predicate usage in WHERE and JOIN clauses for tuning and index advice.", + }, + "pg_repack": { + Name: "pg_repack", Category: "maintenance", + Description: "Rebuilds bloated tables and indexes online with minimal locking.", + }, + "pg_search": { + Name: "pg_search", Category: "search", + Description: "Provides ParadeDB full-text and relevance search features.", + // bm25 is also the access method name used by pg_textsearch; pg_search is the + // canonical provider, so a bm25 index resolves to it. + IndexMethods: []string{"bm25"}, + FunctionPrefixes: []string{"paradedb."}, + }, + "pg_stat_statements": { + Name: "pg_stat_statements", Category: "observability", + Description: "Tracks normalized query execution statistics.", + }, + "pg_textsearch": { + Name: "pg_textsearch", Category: "search", + Description: "Adds BM25-style text search support.", + }, + "pg_trgm": { + Name: "pg_trgm", Category: "text", + Description: "Adds trigram similarity search and fast fuzzy matching indexes.", + OperatorClasses: []string{"gin_trgm_ops", "gist_trgm_ops"}, + Functions: []string{ + "similarity", "word_similarity", "strict_word_similarity", + "show_trgm", "show_limit", "set_limit", + }, + }, + "pgcrypto": { + Name: "pgcrypto", Category: "security", + Description: "Adds hashing, encryption, random bytes, and UUID helpers.", + // gen_random_uuid is deliberately absent: it is built in since PostgreSQL 13. + Functions: []string{ + "crypt", "gen_salt", "gen_random_bytes", "digest", "hmac", + "pgp_sym_encrypt", "pgp_sym_decrypt", "pgp_pub_encrypt", "pgp_pub_decrypt", + "armor", "dearmor", + }, + }, + "pgrouting": { + Name: "pgrouting", Category: "geospatial", + Description: "Adds routing and graph algorithms on top of PostGIS data.", + Requires: []string{"postgis"}, + FunctionPrefixes: []string{"pgr_"}, + }, + "pgstattuple": { + Name: "pgstattuple", Category: "maintenance", + Description: "Reports table and index tuple density and bloat information.", + Functions: []string{"pgstattuple", "pgstatindex", "pgstatginindex", "pg_relpages"}, + }, + "plpython3u": { + Name: "plpython3u", Category: "procedural", + Description: "Lets you write PostgreSQL functions in Python 3.", + }, + "postgis": { + Name: "postgis", Category: "geospatial", + Description: "Adds spatial data types, functions, and indexes.", + IndexMethods: nil, // uses the built-in gist/spgist/brin access methods + OperatorClasses: []string{ + "gist_geometry_ops_2d", "gist_geometry_ops_nd", "gist_geography_ops", + "spgist_geometry_ops_2d", "spgist_geometry_ops_3d", "spgist_geometry_ops_nd", + "brin_geometry_inclusion_ops_2d", "brin_geometry_inclusion_ops_3d", + "brin_geometry_inclusion_ops_4d", "brin_geography_inclusion_ops_2d", + "btree_geometry_ops", "btree_geography_ops", + }, + FunctionPrefixes: []string{"st_"}, + Functions: []string{ + "geometrytype", "addgeometrycolumn", "dropgeometrycolumn", "updategeometrysrid", + "find_srid", "postgis_version", "postgis_full_version", + }, + }, + "postgis_raster": { + Name: "postgis_raster", Category: "geospatial", + Description: "Adds the raster type and raster analysis functions.", + Requires: []string{"postgis"}, + }, + "postgis_topology": { + Name: "postgis_topology", Category: "geospatial", + Description: "Adds topology-aware spatial models and validation tools.", + Requires: []string{"postgis"}, + FunctionPrefixes: []string{"topology."}, + }, + "postgres_fdw": { + Name: "postgres_fdw", Category: "federation", + Description: "Connects PostgreSQL tables to other PostgreSQL servers.", + }, + "timescaledb": { + Name: "timescaledb", Category: "time-series", + Description: "Adds hypertables, compression, retention, and time-series optimizations.", + Functions: []string{ + "create_hypertable", "add_dimension", "time_bucket", "time_bucket_gapfill", + "add_retention_policy", "add_compression_policy", "locf", "interpolate", + }, + }, + "unaccent": { + Name: "unaccent", Category: "text", + Description: "Removes accents and diacritics for normalized text search.", + Functions: []string{"unaccent"}, + }, + "uuid-ossp": { + Name: "uuid-ossp", Category: "utility", + Description: "Generates UUIDs using several algorithms.", + Functions: []string{ + "uuid_generate_v1", "uuid_generate_v1mc", "uuid_generate_v3", + "uuid_generate_v4", "uuid_generate_v5", + "uuid_nil", "uuid_ns_dns", "uuid_ns_url", "uuid_ns_oid", "uuid_ns_x500", + }, + }, + "vector": { + Name: "vector", Category: "ai/search", + Description: "Adds vector data types and similarity search for embeddings.", + IndexMethods: []string{"hnsw", "ivfflat"}, + OperatorClasses: []string{ + "vector_l2_ops", "vector_ip_ops", "vector_cosine_ops", "vector_l1_ops", + "halfvec_l2_ops", "halfvec_ip_ops", "halfvec_cosine_ops", "halfvec_l1_ops", + "sparsevec_l2_ops", "sparsevec_ip_ops", "sparsevec_cosine_ops", "sparsevec_l1_ops", + "bit_hamming_ops", "bit_jaccard_ops", + }, + Functions: []string{"l2_distance", "inner_product", "cosine_distance", "l1_distance", "vector_dims", "vector_norm"}, + }, + "vchord": { + Name: "vchord", Category: "ai/search", + Description: "Adds VectorChord scalable disk-friendly vector indexes compatible with pgvector data types.", + Requires: []string{"vector"}, + IndexMethods: []string{"vchordrq", "vchordg"}, + }, + "ltree": { + Name: "ltree", Category: "document", + Description: "Adds a hierarchical label tree type.", + OperatorClasses: []string{"gist_ltree_ops", "gin_ltree_ops", "gist__ltree_ops"}, + Functions: []string{"subltree", "subpath", "nlevel", "lca", "ltree2text", "text2ltree"}, + }, +} + +// extensionIndexMethods maps an index access method to the extension providing it. +var extensionIndexMethods = buildExtensionIndex(func(ext Extension) []string { return ext.IndexMethods }) + +// extensionOperatorClasses maps an operator class to the extension providing it. +var extensionOperatorClasses = buildExtensionIndex(func(ext Extension) []string { return ext.OperatorClasses }) + +// extensionFunctions maps a function name to the extension providing it. +var extensionFunctions = buildExtensionIndex(func(ext Extension) []string { return ext.Functions }) + +// extensionFunctionPrefixes maps a function name prefix to the extension providing it. +var extensionFunctionPrefixes = buildExtensionIndex(func(ext Extension) []string { return ext.FunctionPrefixes }) + +func buildExtensionIndex(keys func(Extension) []string) map[string]string { + index := make(map[string]string) + for _, ext := range postgresExtensions { + for _, key := range keys(ext) { + // Deterministic on collision: the alphabetically first extension wins. + if existing, ok := index[key]; ok && existing < ext.Name { + continue + } + index[key] = ext.Name + } + } + return index +} + +// LookupExtension returns the registered extension by name. +func LookupExtension(name string) (Extension, bool) { + ext, ok := postgresExtensions[strings.ToLower(strings.TrimSpace(name))] + return ext, ok +} + +// IsKnownExtension reports whether the named extension is registered. +func IsKnownExtension(name string) bool { + _, ok := LookupExtension(name) + return ok +} + +// GetExtensions returns every registered extension name, sorted. +func GetExtensions() []string { + names := make([]string, 0, len(postgresExtensions)) + for name := range postgresExtensions { + names = append(names, name) + } + sort.Strings(names) + return names +} + +// IndexMethodExtension returns the extension providing an index access method +// ("hnsw" -> "vector", "vchordrq" -> "vchord"). Built-in methods return "". +func IndexMethodExtension(method string) string { + return extensionIndexMethods[strings.ToLower(strings.TrimSpace(method))] +} + +// OperatorClassExtension returns the extension providing an operator class +// ("gin_trgm_ops" -> "pg_trgm"). Built-in operator classes return "". +func OperatorClassExtension(opClass string) string { + return extensionOperatorClasses[strings.ToLower(strings.TrimSpace(opClass))] +} + +// ExtensionsForExpression returns the extensions whose functions appear in a SQL +// expression such as a column default, check constraint, index predicate, or view body. +// The result is sorted and deduplicated. +func ExtensionsForExpression(expression string) []string { + if strings.TrimSpace(expression) == "" { + return nil + } + + lower := strings.ToLower(expression) + found := make(map[string]bool) + + for _, call := range sqlFunctionCalls(lower) { + if ext, ok := extensionFunctions[call]; ok { + found[ext] = true + continue + } + for prefix, ext := range extensionFunctionPrefixes { + if strings.HasPrefix(call, prefix) { + found[ext] = true + break + } + } + } + + if len(found) == 0 { + return nil + } + + names := make([]string, 0, len(found)) + for name := range found { + names = append(names, name) + } + sort.Strings(names) + return names +} + +// sqlFunctionCalls returns the lowercase names of every function call in an expression. +// A call is an identifier (optionally schema-qualified) immediately followed by "(". +func sqlFunctionCalls(lowerExpression string) []string { + calls := make([]string, 0, 4) + end := 0 + + for i := 0; i < len(lowerExpression); i++ { + if lowerExpression[i] != '(' { + continue + } + + end = i + // Allow whitespace between the identifier and its opening parenthesis. + for end > 0 && isSQLSpace(lowerExpression[end-1]) { + end-- + } + + start := end + for start > 0 && isSQLIdentifierByte(lowerExpression[start-1]) { + start-- + } + if start == end { + continue + } + // A leading digit means this is not an identifier (e.g. "2("). + if lowerExpression[start] >= '0' && lowerExpression[start] <= '9' { + continue + } + calls = append(calls, lowerExpression[start:end]) + } + + return calls +} + +func isSQLIdentifierByte(b byte) bool { + switch { + case b >= 'a' && b <= 'z', b >= 'A' && b <= 'Z', b >= '0' && b <= '9': + return true + case b == '_', b == '.': + return true + default: + return false + } +} + +func isSQLSpace(b byte) bool { + return b == ' ' || b == '\t' || b == '\n' || b == '\r' +} + +// SortExtensions orders extension names so that dependencies come first (postgis before +// postgis_topology, vector before vchord), with alphabetical order breaking ties. +// Duplicates are removed; unknown names are kept and sorted alphabetically. +func SortExtensions(names []string) []string { + unique := make(map[string]bool, len(names)) + for _, name := range names { + name = strings.ToLower(strings.TrimSpace(name)) + if name != "" { + unique[name] = true + } + } + if len(unique) == 0 { + return nil + } + + pending := make([]string, 0, len(unique)) + for name := range unique { + pending = append(pending, name) + } + sort.Strings(pending) + + sorted := make([]string, 0, len(pending)) + emitted := make(map[string]bool, len(pending)) + + var emit func(name string, seen map[string]bool) + emit = func(name string, seen map[string]bool) { + if emitted[name] || seen[name] { + return + } + seen[name] = true + + if ext, ok := LookupExtension(name); ok { + for _, dependency := range ext.Requires { + // Only order dependencies that are actually being created. + if unique[dependency] { + emit(dependency, seen) + } + } + } + + emitted[name] = true + sorted = append(sorted, name) + } + + for _, name := range pending { + emit(name, make(map[string]bool)) + } + return sorted +} + +// ExtensionDependencies returns the extensions a given extension requires, sorted. +func ExtensionDependencies(name string) []string { + ext, ok := LookupExtension(name) + if !ok || len(ext.Requires) == 0 { + return nil + } + requires := append([]string(nil), ext.Requires...) + sort.Strings(requires) + return requires +} + +// QuoteExtensionName quotes an extension name when it is not a bare SQL identifier, +// e.g. uuid-ossp -> "uuid-ossp". +func QuoteExtensionName(name string) string { + name = strings.TrimSpace(name) + if name == "" { + return "" + } + for i := 0; i < len(name); i++ { + b := name[i] + switch { + case b >= 'a' && b <= 'z', b == '_': + case b >= '0' && b <= '9' && i > 0: + default: + return `"` + strings.ReplaceAll(name, `"`, `""`) + `"` + } + } + return name +} diff --git a/pkg/pgsql/extensions_test.go b/pkg/pgsql/extensions_test.go new file mode 100644 index 0000000..eac8eef --- /dev/null +++ b/pkg/pgsql/extensions_test.go @@ -0,0 +1,165 @@ +package pgsql + +import ( + "reflect" + "strings" + "testing" +) + +func TestExtensionRegistryConsistency(t *testing.T) { + for name, ext := range postgresExtensions { + if name != ext.Name { + t.Errorf("extension registered as %q has Name %q", name, ext.Name) + } + if name != strings.ToLower(name) { + t.Errorf("extension %q must be registered lowercase", name) + } + if ext.Description == "" || ext.Category == "" { + t.Errorf("extension %q is missing a category or description", name) + } + for _, dependency := range ext.Requires { + if !IsKnownExtension(dependency) { + t.Errorf("extension %q requires unregistered extension %q", name, dependency) + } + } + } +} + +// Every extension named by a type in the type registry must itself be registered, +// otherwise a column type would ask for a CREATE EXTENSION nothing knows how to order. +func TestTypeExtensionsAreRegistered(t *testing.T) { + for typeName, spec := range postgresBaseTypes { + if spec.Extension == "" { + continue + } + if !IsKnownExtension(spec.Extension) { + t.Errorf("type %q declares unregistered extension %q", typeName, spec.Extension) + } + } +} + +func TestIndexMethodExtension(t *testing.T) { + tests := map[string]string{ + "hnsw": "vector", + "ivfflat": "vector", + "HNSW": "vector", + "vchordrq": "vchord", + "vchordg": "vchord", + "bm25": "pg_search", + "btree": "", + "gin": "", + "": "", + } + + for method, want := range tests { + if got := IndexMethodExtension(method); got != want { + t.Errorf("IndexMethodExtension(%q) = %q, want %q", method, got, want) + } + } +} + +func TestOperatorClassExtension(t *testing.T) { + tests := map[string]string{ + "gin_trgm_ops": "pg_trgm", + "gist_trgm_ops": "pg_trgm", + "vector_cosine_ops": "vector", + "halfvec_l2_ops": "vector", + "gist_ltree_ops": "ltree", + "gist_geometry_ops_2d": "postgis", + "jsonb_path_ops": "", + "array_ops": "", + "": "", + } + + for opClass, want := range tests { + if got := OperatorClassExtension(opClass); got != want { + t.Errorf("OperatorClassExtension(%q) = %q, want %q", opClass, got, want) + } + } +} + +func TestExtensionsForExpression(t *testing.T) { + tests := []struct { + name string + expression string + want []string + }{ + {"empty", "", nil}, + {"no functions", "status = 'active'", nil}, + {"builtin only", "now()", nil}, + {"uuid-ossp default", "uuid_generate_v4()", []string{"uuid-ossp"}}, + {"gen_random_uuid is builtin", "gen_random_uuid()", nil}, + {"pgcrypto", "crypt(password, gen_salt('bf'))", []string{"pgcrypto"}}, + {"postgis prefix", "ST_Area(geom) > 0", []string{"postgis"}}, + {"paradedb prefix", "paradedb.snippet(body)", []string{"pg_search"}}, + {"jsonschema", "json_matches_schema('{}', payload)", []string{"pg_jsonschema"}}, + {"whitespace before paren", "unaccent ('crème')", []string{"unaccent"}}, + {"multiple sorted", "ST_X(geom) = levenshtein(a, b)::float", []string{"fuzzystrmatch", "postgis"}}, + {"column named like function", "similarity_score > 0.5", nil}, + {"numeric prefix ignored", "2(3)", nil}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := ExtensionsForExpression(tt.expression); !reflect.DeepEqual(got, tt.want) { + t.Errorf("ExtensionsForExpression(%q) = %v, want %v", tt.expression, got, tt.want) + } + }) + } +} + +func TestSortExtensions(t *testing.T) { + tests := []struct { + name string + input []string + want []string + }{ + {"empty", nil, nil}, + {"alphabetical", []string{"pg_trgm", "citext"}, []string{"citext", "pg_trgm"}}, + {"deduplicated", []string{"vector", "vector", " VECTOR "}, []string{"vector"}}, + {"dependency first", []string{"vchord", "vector"}, []string{"vector", "vchord"}}, + { + "postgis dependants", + []string{"postgis_topology", "pgrouting", "postgis"}, + []string{"postgis", "pgrouting", "postgis_topology"}, + }, + {"dependency not requested", []string{"vchord"}, []string{"vchord"}}, + {"unknown names kept", []string{"zzz_custom", "citext"}, []string{"citext", "zzz_custom"}}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := SortExtensions(tt.input); !reflect.DeepEqual(got, tt.want) { + t.Errorf("SortExtensions(%v) = %v, want %v", tt.input, got, tt.want) + } + }) + } +} + +func TestQuoteExtensionName(t *testing.T) { + tests := map[string]string{ + "vector": "vector", + "pg_trgm": "pg_trgm", + "uuid-ossp": `"uuid-ossp"`, + "PostGIS": `"PostGIS"`, + "": "", + } + + for name, want := range tests { + if got := QuoteExtensionName(name); got != want { + t.Errorf("QuoteExtensionName(%q) = %q, want %q", name, got, want) + } + } +} + +func TestExtensionDependencies(t *testing.T) { + if got := ExtensionDependencies("vchord"); !reflect.DeepEqual(got, []string{"vector"}) { + t.Errorf("ExtensionDependencies(vchord) = %v, want [vector]", got) + } + if got := ExtensionDependencies("citext"); got != nil { + t.Errorf("ExtensionDependencies(citext) = %v, want nil", got) + } + if got := ExtensionDependencies("not_an_extension"); got != nil { + t.Errorf("ExtensionDependencies(not_an_extension) = %v, want nil", got) + } +} diff --git a/pkg/pgsql/index_options.go b/pkg/pgsql/index_options.go new file mode 100644 index 0000000..903be69 --- /dev/null +++ b/pkg/pgsql/index_options.go @@ -0,0 +1,248 @@ +package pgsql + +import ( + "strconv" + "strings" +) + +// Index access-method storage parameters, the WITH (...) clause of CREATE INDEX. RelSpec +// carries them through the model in Index.Comment, so the parsing here is deliberately +// strict: only well-formed "key = value" pairs survive, and comment prose is discarded. +// +// Value forms accepted: +// - bare tokens: lists=100, m=16, deduplicate_items=true +// - quoted strings: key_field='id' (pg_search bm25) +// - dollar-quoted blocks: options=$$ [build.internal] lists=[4096] $$ (vchord) + +// ExtractWithClause returns the contents of the first WITH (...) clause in s, without the +// surrounding parentheses. Parentheses inside quoted and dollar-quoted values are ignored, +// so a vchord TOML block survives intact. Returns "" when there is no WITH clause. +func ExtractWithClause(s string) string { + lower := strings.ToLower(s) + + for offset := 0; ; { + idx := strings.Index(lower[offset:], "with") + if idx < 0 { + return "" + } + start := offset + idx + offset = start + 4 + + // "with" must stand as its own word. + if start > 0 && isSQLIdentifierByte(s[start-1]) { + continue + } + + pos := offset + for pos < len(s) && isSQLSpace(s[pos]) { + pos++ + } + if pos >= len(s) || s[pos] != '(' { + continue + } + + if end, ok := matchClosingParen(s, pos); ok { + return s[pos+1 : end] + } + return "" + } +} + +// matchClosingParen returns the index of the ')' matching the '(' at open, skipping over +// quoted and dollar-quoted spans. +func matchClosingParen(s string, open int) (int, bool) { + depth := 0 + for i := open; i < len(s); i++ { + switch s[i] { + case '\'': + end, ok := skipQuoted(s, i) + if !ok { + return 0, false + } + i = end + case '$': + if end, ok := skipDollarQuoted(s, i); ok { + i = end + } + case '(': + depth++ + case ')': + depth-- + if depth == 0 { + return i, true + } + } + } + return 0, false +} + +// skipQuoted returns the index of the closing quote of the single-quoted string starting +// at start, treating ” as an escaped quote. +func skipQuoted(s string, start int) (int, bool) { + for i := start + 1; i < len(s); i++ { + if s[i] != '\'' { + continue + } + if i+1 < len(s) && s[i+1] == '\'' { + i++ + continue + } + return i, true + } + return 0, false +} + +// skipDollarQuoted returns the index of the last byte of the dollar-quoted block starting +// at start ($tag$ … $tag$). Reports false when start does not open one. +func skipDollarQuoted(s string, start int) (int, bool) { + tagEnd := strings.IndexByte(s[start+1:], '$') + if tagEnd < 0 { + return 0, false + } + tag := s[start : start+1+tagEnd+1] + for i := start + 1; i < len(tag); i++ { + if !isSQLIdentifierByte(tag[i]) && tag[i] != '$' { + return 0, false + } + } + + closing := strings.Index(s[start+len(tag):], tag) + if closing < 0 { + return 0, false + } + return start + len(tag) + closing + len(tag) - 1, true +} + +// SplitStorageParameters splits a WITH clause body on top-level commas, leaving quoted and +// dollar-quoted values untouched. +func SplitStorageParameters(clause string) []string { + parts := make([]string, 0, 4) + depth := 0 + start := 0 + + for i := 0; i < len(clause); i++ { + switch clause[i] { + case '\'': + if end, ok := skipQuoted(clause, i); ok { + i = end + } + case '$': + if end, ok := skipDollarQuoted(clause, i); ok { + i = end + } + case '(', '[': + depth++ + case ')', ']': + depth-- + case ',': + if depth == 0 { + parts = append(parts, clause[start:i]) + start = i + 1 + } + } + } + parts = append(parts, clause[start:]) + + trimmed := make([]string, 0, len(parts)) + for _, part := range parts { + if part = strings.TrimSpace(part); part != "" { + trimmed = append(trimmed, part) + } + } + return trimmed +} + +// ParseStorageParameter splits one "key = value" storage parameter. It reports false for +// anything that is not a well-formed parameter, which is how comment prose is filtered out. +func ParseStorageParameter(part string) (key, value string, ok bool) { + key, value, found := strings.Cut(part, "=") + if !found { + return "", "", false + } + + key = strings.ToLower(strings.TrimSpace(key)) + value = strings.TrimSpace(value) + if key == "" || value == "" || !isBareIdentifier(key) { + return "", "", false + } + if !isStorageParameterValue(value) { + return "", "", false + } + return key, value, true +} + +// FormatStorageParameters renders a WITH clause body as a canonical "key = value" list, +// dropping anything malformed. Returns "" when nothing survives. +func FormatStorageParameters(clause string) string { + params := make([]string, 0, 4) + for _, part := range SplitStorageParameters(clause) { + key, value, ok := ParseStorageParameter(part) + if !ok { + continue + } + params = append(params, key+" = "+value) + } + return strings.Join(params, ", ") +} + +// NormalizeStorageParameterValue unquotes a value that PostgreSQL rendered as a string but +// that is really a number, so that pg_indexes output (lists='100') and hand-written models +// (lists=100) normalize identically. Non-numeric quoted values keep their quotes because +// some access methods require a string (pg_search's key_field='id'). +func NormalizeStorageParameterValue(value string) string { + value = strings.TrimSpace(value) + if len(value) < 2 || value[0] != '\'' || value[len(value)-1] != '\'' { + return value + } + + inner := strings.ReplaceAll(value[1:len(value)-1], "''", "'") + if _, err := strconv.ParseFloat(inner, 64); err == nil { + return inner + } + if strings.EqualFold(inner, "true") || strings.EqualFold(inner, "false") { + return strings.ToLower(inner) + } + return value +} + +func isBareIdentifier(s string) bool { + if s == "" { + return false + } + for i := 0; i < len(s); i++ { + b := s[i] + switch { + case b >= 'a' && b <= 'z', b >= 'A' && b <= 'Z', b == '_': + case b >= '0' && b <= '9' && i > 0: + default: + return false + } + } + return true +} + +// isStorageParameterValue reports whether value is a bare token, a complete quoted string, +// or a complete dollar-quoted block. +func isStorageParameterValue(value string) bool { + switch { + case value == "": + return false + case value[0] == '\'': + end, ok := skipQuoted(value, 0) + return ok && end == len(value)-1 + case value[0] == '$': + end, ok := skipDollarQuoted(value, 0) + return ok && end == len(value)-1 + } + + for i := 0; i < len(value); i++ { + b := value[i] + switch { + case b >= 'a' && b <= 'z', b >= 'A' && b <= 'Z', b >= '0' && b <= '9': + case b == '_', b == '.', b == '-', b == '+': + default: + return false + } + } + return true +} diff --git a/pkg/pgsql/index_options_test.go b/pkg/pgsql/index_options_test.go new file mode 100644 index 0000000..2c6c10a --- /dev/null +++ b/pkg/pgsql/index_options_test.go @@ -0,0 +1,111 @@ +package pgsql + +import ( + "reflect" + "testing" +) + +func TestExtractWithClause(t *testing.T) { + tests := []struct { + name string + input string + want string + }{ + {"empty", "", ""}, + {"no clause", "opclass=vector_cosine_ops", ""}, + {"simple", "WITH (lists=100)", "lists=100"}, + {"lowercase", "with (m = 16, ef_construction = 64)", "m = 16, ef_construction = 64"}, + { + "index definition", + "CREATE INDEX i ON t USING ivfflat (embedding vector_cosine_ops) WITH (lists='100')", + "lists='100'", + }, + {"paren inside quotes", "with (key_field='id(x)')", "key_field='id(x)'"}, + {"dollar quoted", "with (options = $$f(x)$$)", "options = $$f(x)$$"}, + {"word boundary", "swith (lists=100)", ""}, + {"not followed by paren", "with lists=100", ""}, + {"unterminated", "with (lists=100", ""}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := ExtractWithClause(tt.input); got != tt.want { + t.Errorf("ExtractWithClause(%q) = %q, want %q", tt.input, got, tt.want) + } + }) + } +} + +func TestSplitStorageParameters(t *testing.T) { + tests := []struct { + name string + input string + want []string + }{ + {"empty", "", []string{}}, + {"single", "lists=100", []string{"lists=100"}}, + {"multiple", "m = 16, ef_construction = 64", []string{"m = 16", "ef_construction = 64"}}, + {"comma in quotes", "key_field='a,b', m=16", []string{"key_field='a,b'", "m=16"}}, + {"comma in dollar quotes", "options=$$a,b$$, m=16", []string{"options=$$a,b$$", "m=16"}}, + {"comma in brackets", "options=[1,2], m=16", []string{"options=[1,2]", "m=16"}}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := SplitStorageParameters(tt.input); !reflect.DeepEqual(got, tt.want) { + t.Errorf("SplitStorageParameters(%q) = %v, want %v", tt.input, got, tt.want) + } + }) + } +} + +func TestParseStorageParameter(t *testing.T) { + tests := []struct { + name string + input string + wantKey string + wantValue string + wantOK bool + }{ + {"bare", "lists=100", "lists", "100", true}, + {"spaced and uppercased key", " Lists = 100 ", "lists", "100", true}, + {"quoted", "key_field='id'", "key_field", "'id'", true}, + {"dollar quoted", "options=$$a$$", "options", "$$a$$", true}, + {"boolean", "deduplicate_items=true", "deduplicate_items", "true", true}, + {"float", "fillfactor=90.5", "fillfactor", "90.5", true}, + {"no equals", "please drop everything", "", "", false}, + {"empty value", "lists=", "", "", false}, + {"quoted key rejected", "'lists'=100", "", "", false}, + {"injection rejected", "lists=100); drop table t", "", "", false}, + {"unterminated quote rejected", "key_field='id", "", "", false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + key, value, ok := ParseStorageParameter(tt.input) + if key != tt.wantKey || value != tt.wantValue || ok != tt.wantOK { + t.Errorf("ParseStorageParameter(%q) = (%q, %q, %v), want (%q, %q, %v)", + tt.input, key, value, ok, tt.wantKey, tt.wantValue, tt.wantOK) + } + }) + } +} + +func TestNormalizeStorageParameterValue(t *testing.T) { + tests := map[string]string{ + "'100'": "100", + "'90.5'": "90.5", + "'true'": "true", + "'id'": "'id'", + "100": "100", + "$$a,b$$": "$$a,b$$", + "'": "'", + "''": "''", + } + + for input, want := range tests { + if got := NormalizeStorageParameterValue(input); got != want { + t.Errorf("NormalizeStorageParameterValue(%q) = %q, want %q", input, got, want) + } + } +} diff --git a/pkg/pgsql/types_registry.go b/pkg/pgsql/types_registry.go index b6a14e8..9fa1695 100644 --- a/pkg/pgsql/types_registry.go +++ b/pkg/pgsql/types_registry.go @@ -2,6 +2,7 @@ package pgsql import ( "sort" + "strconv" "strings" ) @@ -9,6 +10,14 @@ import ( type TypeSpec struct { SupportsLength bool SupportsPrecision bool + + // SupportsTypeModifier marks types whose "(...)" modifier is opaque and must be + // preserved verbatim (e.g. vector(1536), geometry(Point,4326)) instead of being + // decomposed into Length/Precision/Scale. + SupportsTypeModifier bool + + // Extension is the PostgreSQL extension providing the type; empty for built-ins. + Extension string } var postgresBaseTypes = map[string]TypeSpec{ @@ -104,14 +113,28 @@ var postgresBaseTypes = map[string]TypeSpec{ "void": {}, // Common extensions - "citext": {}, - "hstore": {}, - "ltree": {}, - "lquery": {}, - "ltxtquery": {}, - "vector": {}, // pgvector: keep explicit modifier form (vector(dim)) - "halfvec": {}, // pgvector: keep explicit modifier form (halfvec(dim)) - "sparsevec": {}, // pgvector: keep explicit modifier form (sparsevec(dim)) + "citext": {Extension: "citext"}, + "hstore": {Extension: "hstore"}, + "ltree": {Extension: "ltree"}, + "lquery": {Extension: "ltree"}, + "ltxtquery": {Extension: "ltree"}, + + // pgvector: modifier form is opaque (vector(dim), sparsevec(dim)) + "vector": {SupportsTypeModifier: true, Extension: "vector"}, + "halfvec": {SupportsTypeModifier: true, Extension: "vector"}, + "sparsevec": {SupportsTypeModifier: true, Extension: "vector"}, + + // PostGIS: geometry/geography carry an opaque modifier (geometry(PointZ,4326)) + "geometry": {SupportsTypeModifier: true, Extension: "postgis"}, + "geography": {SupportsTypeModifier: true, Extension: "postgis"}, + "box2d": {Extension: "postgis"}, + "box3d": {Extension: "postgis"}, + "geometry_dump": {Extension: "postgis"}, + "geomval": {Extension: "postgis"}, + "spheroid": {Extension: "postgis"}, + "valid_detail": {Extension: "postgis"}, + "raster": {SupportsTypeModifier: true, Extension: "postgis_raster"}, + "topogeometry": {Extension: "postgis_topology"}, } var postgresTypeAliases = map[string]string{ @@ -346,3 +369,72 @@ func stripArraySuffixes(t string) string { func normalizeTypeToken(t string) string { return strings.Join(strings.Fields(strings.TrimSpace(t)), " ") } + +// SupportsTypeModifier reports if this SQL type carries an opaque "(...)" modifier +// that must be preserved verbatim (e.g. vector(1536), geometry(Point,4326)). +func SupportsTypeModifier(sqlType string) bool { + base := CanonicalizeBaseType(ExtractBaseTypeLower(sqlType)) + spec, ok := postgresBaseTypes[base] + return ok && spec.SupportsTypeModifier +} + +// TypeExtension returns the PostgreSQL extension providing the given type +// ("postgis", "vector", "citext", …). Built-in types return "". +func TypeExtension(sqlType string) string { + base := CanonicalizeBaseType(ExtractBaseTypeLower(sqlType)) + return postgresBaseTypes[base].Extension +} + +// IsSpatialType reports whether the type comes from PostGIS (geometry, geography, +// raster, topogeometry, …). +func IsSpatialType(sqlType string) bool { + return strings.HasPrefix(TypeExtension(sqlType), "postgis") +} + +// IsVectorType reports whether the type comes from pgvector (vector, halfvec, sparsevec). +func IsVectorType(sqlType string) bool { + return TypeExtension(sqlType) == "vector" +} + +// TypeModifier returns the raw "(...)" modifier of a SQL type without the parentheses, +// or "" when the type has none. Array suffixes are ignored. +// Example: geometry(PointZ,4326)[] -> "PointZ,4326". +func TypeModifier(sqlType string) string { + t := stripArraySuffixes(normalizeTypeToken(sqlType)) + start := strings.Index(t, "(") + end := strings.LastIndex(t, ")") + if start < 0 || end < start { + return "" + } + return strings.TrimSpace(t[start+1 : end]) +} + +// SpatialSRID returns the SRID declared in a PostGIS type modifier, or 0 when absent. +// Example: geometry(Point,4326) -> 4326. +func SpatialSRID(sqlType string) int { + if !IsSpatialType(sqlType) { + return 0 + } + parts := strings.Split(TypeModifier(sqlType), ",") + if len(parts) < 2 { + return 0 + } + srid, err := strconv.Atoi(strings.TrimSpace(parts[len(parts)-1])) + if err != nil { + return 0 + } + return srid +} + +// SpatialGeometryType returns the geometry subtype declared in a PostGIS type modifier +// ("Point", "MultiPolygonZ", …), or "" when absent. +func SpatialGeometryType(sqlType string) string { + if !IsSpatialType(sqlType) { + return "" + } + modifier := TypeModifier(sqlType) + if modifier == "" { + return "" + } + return strings.TrimSpace(strings.Split(modifier, ",")[0]) +} diff --git a/pkg/pgsql/types_registry_test.go b/pkg/pgsql/types_registry_test.go index 6653bbf..5087cb4 100644 --- a/pkg/pgsql/types_registry_test.go +++ b/pkg/pgsql/types_registry_test.go @@ -145,3 +145,104 @@ func TestEquivalentSQLTypeVariants(t *testing.T) { }) } } + +func TestExtensionTypes(t *testing.T) { + tests := []struct { + name string + sqlType string + wantKnown bool + wantExtension string + wantSpatial bool + wantVector bool + wantModifier bool + }{ + {"geometry with modifier", "geometry(Point,4326)", true, "postgis", true, false, true}, + {"geography", "geography", true, "postgis", true, false, true}, + {"geometry array", "geometry[]", true, "postgis", true, false, true}, + {"box2d", "box2d", true, "postgis", true, false, false}, + {"raster", "raster", true, "postgis_raster", true, false, true}, + {"topogeometry", "topogeometry", true, "postgis_topology", true, false, false}, + {"vector", "vector(1536)", true, "vector", false, true, true}, + {"halfvec", "halfvec(768)", true, "vector", false, true, true}, + {"sparsevec", "sparsevec(1000)", true, "vector", false, true, true}, + {"citext", "citext", true, "citext", false, false, false}, + {"builtin text", "text", true, "", false, false, false}, + {"builtin point is not postgis", "point", true, "", false, false, false}, + {"unknown type", "mytype", false, "", false, false, false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := IsKnownPostgresType(tt.sqlType); got != tt.wantKnown { + t.Errorf("IsKnownPostgresType(%q) = %v, want %v", tt.sqlType, got, tt.wantKnown) + } + if got := TypeExtension(tt.sqlType); got != tt.wantExtension { + t.Errorf("TypeExtension(%q) = %q, want %q", tt.sqlType, got, tt.wantExtension) + } + if got := IsSpatialType(tt.sqlType); got != tt.wantSpatial { + t.Errorf("IsSpatialType(%q) = %v, want %v", tt.sqlType, got, tt.wantSpatial) + } + if got := IsVectorType(tt.sqlType); got != tt.wantVector { + t.Errorf("IsVectorType(%q) = %v, want %v", tt.sqlType, got, tt.wantVector) + } + if got := SupportsTypeModifier(tt.sqlType); got != tt.wantModifier { + t.Errorf("SupportsTypeModifier(%q) = %v, want %v", tt.sqlType, got, tt.wantModifier) + } + }) + } +} + +func TestExtensionTypesDoNotSupportLengthOrPrecision(t *testing.T) { + for _, sqlType := range []string{"geometry(Point,4326)", "geography", "vector(1536)", "halfvec(768)"} { + if SupportsLength(sqlType) { + t.Errorf("SupportsLength(%q) = true, want false", sqlType) + } + if SupportsPrecision(sqlType) { + t.Errorf("SupportsPrecision(%q) = true, want false", sqlType) + } + } +} + +func TestSpatialTypeModifier(t *testing.T) { + tests := []struct { + sqlType string + wantModifier string + wantGeomType string + wantSRID int + }{ + {"geometry(Point,4326)", "Point,4326", "Point", 4326}, + {"geometry(MultiPolygonZ, 3857)", "MultiPolygonZ, 3857", "MultiPolygonZ", 3857}, + {"geography(Point)", "Point", "Point", 0}, + {"geometry", "", "", 0}, + {"geometry(Point,4326)[]", "Point,4326", "Point", 4326}, + {"vector(1536)", "1536", "", 0}, + } + + for _, tt := range tests { + t.Run(tt.sqlType, func(t *testing.T) { + if got := TypeModifier(tt.sqlType); got != tt.wantModifier { + t.Errorf("TypeModifier() = %q, want %q", got, tt.wantModifier) + } + if got := SpatialGeometryType(tt.sqlType); got != tt.wantGeomType { + t.Errorf("SpatialGeometryType() = %q, want %q", got, tt.wantGeomType) + } + if got := SpatialSRID(tt.sqlType); got != tt.wantSRID { + t.Errorf("SpatialSRID() = %d, want %d", got, tt.wantSRID) + } + }) + } +} + +func TestNormalizeEquivalentSQLTypePreservesExtensionModifiers(t *testing.T) { + tests := map[string]string{ + "geometry(Point,4326)": "geometry(Point,4326)", + "vector(1536)": "vector(1536)", + "geography(Point)[]": "geography(Point)[]", + } + + for input, want := range tests { + if got := NormalizeEquivalentSQLType(input); got != want { + t.Errorf("NormalizeEquivalentSQLType(%q) = %q, want %q", input, got, want) + } + } +} diff --git a/pkg/readers/dbml/reader_test.go b/pkg/readers/dbml/reader_test.go index 463bdaa..4c76c72 100644 --- a/pkg/readers/dbml/reader_test.go +++ b/pkg/readers/dbml/reader_test.go @@ -863,6 +863,13 @@ func TestParseColumn_PostgresTypes(t *testing.T) { wantName: "embedding", wantType: "vector(1536)", }, + { + name: "postgis geometry with type modifier", + line: "location geometry(Point,4326) [not null]", + wantName: "location", + wantType: "geometry(Point,4326)", + wantNotNull: true, + }, { name: "multi word timestamp type", line: "published_at timestamp with time zone", diff --git a/pkg/readers/pgsql/README.md b/pkg/readers/pgsql/README.md index 78de0c4..d25f0eb 100644 --- a/pkg/readers/pgsql/README.md +++ b/pkg/readers/pgsql/README.md @@ -128,6 +128,27 @@ sessions so they are identifiable in `pg_stat_activity`. If you provide - Sequence properties - Associated tables +## Extension Types (PostGIS, pgvector) + +- Extension column types keep their catalog-formatted form: `geometry(Point,4326)`, + `geography(Point)`, `vector(1536)`, `halfvec(768)`, `citext`, arrays included. +- Built-in types are canonicalized and their dimensions moved to + `Column.Length` / `Precision` / `Scale`; extension modifiers stay in `Column.Type`. +- Index access methods are read from the definition as-is: `gist`, `spgist`, `brin`, `hnsw`, + `ivfflat`, `vchordrq`, `vchordg`, `bm25`. +- Operator class and `WITH (...)` parameters have no model field, so they are stored in + `Index.Comment` in the form the PostgreSQL writer reads back: + + ``` + opclass=vector_cosine_ops; with (m=16, ef_construction=64) + ``` + + Ordering modifiers (`DESC`, `NULLS LAST`, `COLLATE`) are not treated as operator classes. + Numeric parameter values are unquoted (`lists='100'` -> `lists=100`); string values keep + their quotes (`key_field='id'`), and dollar-quoted values are preserved whole. +- Installed extensions are read from `pg_extension` into `schema.Metadata["extensions"]` + (only extensions RelSpec recognizes), so a read/write round-trip re-creates them. + ## Notes - Requires PostgreSQL connection permissions diff --git a/pkg/readers/pgsql/queries.go b/pkg/readers/pgsql/queries.go index e2a54d2..b8809fa 100644 --- a/pkg/readers/pgsql/queries.go +++ b/pkg/readers/pgsql/queries.go @@ -5,6 +5,7 @@ import ( "strings" "git.warky.dev/wdevs/relspecgo/pkg/models" + "git.warky.dev/wdevs/relspecgo/pkg/pgsql" ) // querySchemas retrieves all non-system schemas from the database @@ -46,6 +47,41 @@ func (r *Reader) querySchemas() ([]*models.Schema, error) { return schemas, rows.Err() } +// queryExtensions retrieves the extensions installed into a schema. Only extensions RelSpec +// recognizes are kept, so a round-trip never emits a CREATE EXTENSION the writer cannot +// order; plpgsql is not registered and is therefore skipped along with other built-ins. +func (r *Reader) queryExtensions(schemaName string) ([]string, error) { + query := ` + SELECT e.extname + FROM pg_extension e + JOIN pg_namespace n ON n.oid = e.extnamespace + WHERE n.nspname = $1 + ORDER BY e.extname + ` + + rows, err := r.conn.Query(r.ctx, query, schemaName) + if err != nil { + return nil, err + } + defer rows.Close() + + extensions := make([]string, 0) + for rows.Next() { + var name string + if err := rows.Scan(&name); err != nil { + return nil, err + } + if pgsql.IsKnownExtension(name) { + extensions = append(extensions, name) + } + } + if err := rows.Err(); err != nil { + return nil, err + } + + return pgsql.SortExtensions(extensions), nil +} + // queryTables retrieves all tables for a given schema func (r *Reader) queryTables(schemaName string) ([]*models.Table, error) { query := ` @@ -597,6 +633,7 @@ func (r *Reader) parseIndexDefinition(indexName, tableName, schema, indexDef str } // Extract columns - pattern: (column1, column2, ...) + opClass := "" columnsRegex := regexp.MustCompile(`\(([^)]+)\)`) if matches := columnsRegex.FindStringSubmatch(indexDef); len(matches) > 1 { columnsStr := matches[1] @@ -604,8 +641,17 @@ func (r *Reader) parseIndexDefinition(indexName, tableName, schema, indexDef str columnParts := strings.Split(columnsStr, ",") for _, col := range columnParts { col = strings.TrimSpace(col) + fields := strings.Fields(col) + if len(fields) == 0 { + continue + } + // Remember an explicit operator class (e.g. "embedding vector_cosine_ops") + // so the writer can reproduce it; ordering modifiers are not operator classes. + if opClass == "" && len(fields) > 1 { + opClass = extractIndexOperatorClass(fields[1:]) + } // Remove any ordering (ASC/DESC) or other modifiers - col = strings.Fields(col)[0] + col = fields[0] // Remove parentheses if it's an expression if !strings.Contains(col, "(") { index.Columns = append(index.Columns, col) @@ -613,6 +659,15 @@ func (r *Reader) parseIndexDefinition(indexName, tableName, schema, indexDef str } } + // Extract access method storage parameters, e.g. WITH (lists='100') + storageParams := normalizeIndexStorageParams(pgsql.ExtractWithClause(indexDef)) + + // Operator class and storage parameters have no dedicated model fields; carry them in + // the comment hint the PostgreSQL writer reads back. + if hint := buildIndexHint(opClass, storageParams); hint != "" && index.Comment == "" { + index.Comment = hint + } + // Extract WHERE clause for partial indexes whereRegex := regexp.MustCompile(`WHERE\s+(.+)$`) if matches := whereRegex.FindStringSubmatch(indexDef); len(matches) > 1 { @@ -622,6 +677,52 @@ func (r *Reader) parseIndexDefinition(indexName, tableName, schema, indexDef str return index, nil } +// indexOrderingKeywords are column modifiers that are not operator classes. +var indexOrderingKeywords = map[string]bool{ + "asc": true, "desc": true, "nulls": true, "first": true, "last": true, "collate": true, +} + +// extractIndexOperatorClass picks the operator class out of a column's trailing modifiers. +// Returns "" when the modifiers are only ordering keywords. +func extractIndexOperatorClass(modifiers []string) string { + for _, modifier := range modifiers { + lower := strings.ToLower(strings.TrimSpace(modifier)) + if lower == "" || indexOrderingKeywords[lower] { + continue + } + return lower + } + return "" +} + +// normalizeIndexStorageParams rewrites "m='16', ef_construction='64'" as "m=16, +// ef_construction=64". Non-numeric values keep their quotes because some access methods +// require a string literal (pg_search's key_field='id'). +func normalizeIndexStorageParams(params string) string { + normalized := make([]string, 0, 4) + for _, part := range pgsql.SplitStorageParameters(params) { + key, value, ok := pgsql.ParseStorageParameter(part) + if !ok { + continue + } + normalized = append(normalized, key+"="+pgsql.NormalizeStorageParameterValue(value)) + } + return strings.Join(normalized, ", ") +} + +// buildIndexHint renders the operator class and storage parameters in the form the +// PostgreSQL writer parses back out of an index comment. +func buildIndexHint(opClass, storageParams string) string { + parts := make([]string, 0, 2) + if opClass != "" { + parts = append(parts, "opclass="+opClass) + } + if storageParams != "" { + parts = append(parts, "with ("+storageParams+")") + } + return strings.Join(parts, "; ") +} + // normalizePostgresDefault converts a raw PostgreSQL column_default expression into the // unquoted string value that the model convention expects. PostgreSQL stores string // literal defaults as 'value' or 'value'::type (e.g. '{}'::text[]), while every other diff --git a/pkg/readers/pgsql/reader.go b/pkg/readers/pgsql/reader.go index 5c14966..1adcd16 100644 --- a/pkg/readers/pgsql/reader.go +++ b/pkg/readers/pgsql/reader.go @@ -88,6 +88,18 @@ func (r *Reader) ReadDatabase() (*models.Database, error) { } schema.Sequences = sequences + // Query extensions installed into this schema + extensions, err := r.queryExtensions(schema.Name) + if err != nil { + return nil, fmt.Errorf("failed to query extensions for schema %s: %w", schema.Name, err) + } + if len(extensions) > 0 { + if schema.Metadata == nil { + schema.Metadata = make(map[string]any) + } + schema.Metadata["extensions"] = extensions + } + // Query columns for tables and views columnsMap, err := r.queryColumns(schema.Name) if err != nil { @@ -278,11 +290,6 @@ func (r *Reader) mapDataType(pgType, udtName, formattedType string, hasNextval b } } - // information_schema reports arrays generically as "ARRAY" with udt_name like "_text". - if strings.EqualFold(pgType, "ARRAY") && strings.HasPrefix(udtName, "_") && len(udtName) > 1 { - return udtName[1:] + "[]" - } - // Use the database-formatted type when available. For known built-in types, strip // embedded dimensions (they are stored in column.Length/Precision/Scale separately). // For unknown/custom types, keep the full formatted string (e.g. vector(1536)). @@ -303,6 +310,13 @@ func (r *Reader) mapDataType(pgType, udtName, formattedType string, hasNextval b return formattedType } + // information_schema reports arrays generically as "ARRAY" with udt_name like "_text". + // Only reached when the catalog-formatted type is unavailable, which is the one case + // where the element modifier (e.g. geometry(Point,4326)[]) cannot be recovered. + if strings.EqualFold(pgType, "ARRAY") && strings.HasPrefix(udtName, "_") && len(udtName) > 1 { + return udtName[1:] + "[]" + } + // Fall back to normalizing the information_schema type name directly. canonical := pgsql.NormalizePGType(normalizedPGType) if pgsql.IsKnownPGBaseType(canonical) { diff --git a/pkg/readers/pgsql/reader_test.go b/pkg/readers/pgsql/reader_test.go index 022a892..142ad4e 100644 --- a/pkg/readers/pgsql/reader_test.go +++ b/pkg/readers/pgsql/reader_test.go @@ -392,3 +392,101 @@ func BenchmarkReader_ReadDatabase(b *testing.B) { } } } + +func TestParseIndexDefinition_ExtensionIndexes(t *testing.T) { + reader := &Reader{} + + tests := []struct { + name string + indexDef string + wantType string + wantColumns []string + wantComment string + }{ + { + name: "hnsw vector index with storage parameters", + indexDef: "CREATE INDEX idx_docs_embedding ON public.docs USING hnsw (embedding vector_cosine_ops) WITH (m='16', ef_construction='64')", + wantType: "hnsw", + wantColumns: []string{"embedding"}, + wantComment: "opclass=vector_cosine_ops; with (m=16, ef_construction=64)", + }, + { + name: "ivfflat vector index", + indexDef: "CREATE INDEX idx_docs_embedding ON public.docs USING ivfflat (embedding vector_l2_ops) WITH (lists='100')", + wantType: "ivfflat", + wantColumns: []string{"embedding"}, + wantComment: "opclass=vector_l2_ops; with (lists=100)", + }, + { + name: "gist geometry index with default operator class", + indexDef: "CREATE INDEX idx_places_geom ON public.places USING gist (geom)", + wantType: "gist", + wantColumns: []string{"geom"}, + wantComment: "", + }, + { + name: "gist geometry index with explicit operator class", + indexDef: "CREATE INDEX idx_places_geom ON public.places USING gist (geom gist_geometry_ops_nd)", + wantType: "gist", + wantColumns: []string{"geom"}, + wantComment: "opclass=gist_geometry_ops_nd", + }, + { + name: "btree ordering modifiers are not operator classes", + indexDef: "CREATE INDEX idx_users_created ON public.users USING btree (created_at DESC NULLS LAST)", + wantType: "btree", + wantColumns: []string{"created_at"}, + wantComment: "", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + index, err := reader.parseIndexDefinition("idx", "tbl", "public", tt.indexDef) + if err != nil { + t.Fatalf("parseIndexDefinition() error = %v", err) + } + + if index.Type != tt.wantType { + t.Errorf("Type = %q, want %q", index.Type, tt.wantType) + } + if len(index.Columns) != len(tt.wantColumns) { + t.Fatalf("Columns = %v, want %v", index.Columns, tt.wantColumns) + } + for i, col := range tt.wantColumns { + if index.Columns[i] != col { + t.Errorf("Columns[%d] = %q, want %q", i, index.Columns[i], col) + } + } + if index.Comment != tt.wantComment { + t.Errorf("Comment = %q, want %q", index.Comment, tt.wantComment) + } + }) + } +} + +func TestMapDataType_ExtensionTypesPreserveModifiers(t *testing.T) { + reader := &Reader{} + + tests := []struct { + name string + pgType string + udtName string + formattedType string + want string + }{ + {"postgis geometry", "USER-DEFINED", "geometry", "geometry(Point,4326)", "geometry(Point,4326)"}, + {"postgis geography", "USER-DEFINED", "geography", "geography(Point,4326)", "geography(Point,4326)"}, + {"postgis geometry without modifier", "USER-DEFINED", "geometry", "geometry", "geometry"}, + {"pgvector halfvec", "USER-DEFINED", "halfvec", "halfvec(768)", "halfvec(768)"}, + {"postgis geometry array", "ARRAY", "_geometry", "geometry(Point,4326)[]", "geometry(Point,4326)[]"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := reader.mapDataType(tt.pgType, tt.udtName, tt.formattedType, false); got != tt.want { + t.Errorf("mapDataType() = %q, want %q", got, tt.want) + } + }) + } +} diff --git a/pkg/writers/pgsql/README.md b/pkg/writers/pgsql/README.md index 5377b66..64dd2ff 100644 --- a/pkg/writers/pgsql/README.md +++ b/pkg/writers/pgsql/README.md @@ -171,6 +171,7 @@ When `include_audit` is enabled, adds: - Function-based indexes - Concurrent index creation (`CREATE INDEX CONCURRENTLY`) via `Index.Concurrent` - Check constraints with expressions +- Extension types and indexes: PostGIS, pgvector, citext, hstore, ltree (see below) ## Data Types @@ -186,6 +187,108 @@ Supports all PostgreSQL data types: - Network: INET, CIDR, MACADDR - Special: ARRAY, HSTORE +## Extension Types (PostGIS, pgvector) + +Extension column types are preserved verbatim, including their type modifier: + +| Type | Example column type | Extension | +|------|---------------------|-----------| +| PostGIS | `geometry(Point,4326)`, `geography(Point)`, `box2d`, `raster` | `postgis`, `postgis_raster`, `postgis_topology` | +| pgvector | `vector(1536)`, `halfvec(768)`, `sparsevec(1000)` | `vector` | +| Other | `citext`, `hstore`, `ltree` | `citext`, `hstore`, `ltree` | + +`CREATE EXTENSION IF NOT EXISTS ;` is emitted automatically for every extension the +schema needs. See [Extensions](#extensions). + +### Extension Indexes + +`Index.Type` selects the access method: `gist`, `spgist`, `brin` (PostGIS), `hnsw`, `ivfflat` +(pgvector), `vchordrq`, `vchordg` (VectorChord), `bm25` (pg_search). + +Operator class and access-method parameters ride in `Index.Comment`: + +``` +opclass=vector_l2_ops; with (lists=100) +``` + +- `opclass=` — used only when compatible with the column type; otherwise ignored. + Bare operator class names in the comment (e.g. `gin_trgm_ops`) are also recognized. +- `with (k=v, …)` — rendered as `WITH (k = v, …)`. Only well-formed `key = value` pairs are + kept, so comment prose never reaches the DDL. Values may be bare (`lists=100`), quoted + (`key_field='id'`), or dollar-quoted (`options=$$[build.internal]$$`). + +Defaults when no operator class is requested: + +| Access method | Column type | Emitted operator class | +|---------------|-------------|------------------------| +| `hnsw`, `ivfflat`, `vchordrq`, `vchordg` | `vector` / `halfvec` / `sparsevec` / `bit` | `vector_cosine_ops` / `halfvec_cosine_ops` / `sparsevec_cosine_ops` / `bit_hamming_ops` | +| `gist`, `spgist`, `brin` | `geometry`, `geography` | none (PostGIS default operator class) | +| `gin` | text / `jsonb` / array | `gin_trgm_ops` / `jsonb_ops` / `array_ops` | + +pgvector defines no default operator class, so a vector index always names one. + +```sql +CREATE INDEX IF NOT EXISTS idx_documents_embedding + ON public.documents USING ivfflat (embedding vector_cosine_ops) WITH (lists = 100); +CREATE INDEX IF NOT EXISTS idx_documents_location + ON public.documents USING gist (location); +``` + +Migrations only recreate an index when both sides specify a hint and they differ, so a model +without hints does not churn against a live database. + +## Extensions + +`CREATE EXTENSION IF NOT EXISTS ;` is emitted per schema, deduplicated and ordered so +dependencies come first (`postgis` before `postgis_topology`/`postgis_raster`/`pgrouting`, +`vector` before `vchord`). Names needing quoting are quoted: `CREATE EXTENSION IF NOT EXISTS "uuid-ossp";` + +### Detection + +| Source | Example | Extension | +|--------|---------|-----------| +| Column type | `vector(1536)`, `geometry(Point,4326)`, `citext`, `ltree` | `vector`, `postgis`, `citext`, `ltree` | +| Index access method | `hnsw`, `ivfflat` / `vchordrq`, `vchordg` / `bm25` | `vector` / `vchord` / `pg_search` | +| Operator class | `gin_trgm_ops`, `gist_ltree_ops` | `pg_trgm`, `ltree` | +| GIN/GiST on a scalar type | `USING gin (views)` | `btree_gin` / `btree_gist` | +| Function in a default, CHECK, index `WHERE`, or view body | `uuid_generate_v4()`, `crypt()`, `ST_Area()`, `unaccent()`, `json_matches_schema()` | `uuid-ossp`, `pgcrypto`, `postgis`, `unaccent`, `pg_jsonschema` | + +`gen_random_uuid()` is built in since PostgreSQL 13 and does not pull in `pgcrypto`. + +### Declaring extensions explicitly + +Extensions that leave no trace in the schema go in `schema.Metadata["extensions"]`, as a list +or a comma-separated string. Dependencies are pulled in automatically; unknown names are kept +as given. The PostgreSQL reader populates this from `pg_extension` for the schemas it reads. + +```yaml +metadata: + extensions: [pg_cron, timescaledb, pg_stat_statements] +``` + +### Recognized extensions + +| Category | Extensions | +|----------|------------| +| ai/search | `vector`, `vchord` | +| document | `hstore`, `ltree` | +| federation | `postgres_fdw` | +| geospatial | `postgis`, `postgis_raster`, `postgis_topology`, `pgrouting` | +| indexing | `btree_gin`, `btree_gist` | +| integration | `http` | +| integrity | `amcheck` | +| jobs / scheduling | `pg_background`, `pg_cron` | +| maintenance | `pg_repack`, `pgstattuple` | +| observability | `pg_qualstats`, `pg_stat_statements` | +| partitioning | `pg_partman` | +| procedural | `plpython3u` | +| search | `pg_search`, `pg_textsearch` | +| security | `pgcrypto` | +| text | `citext`, `fuzzystrmatch`, `pg_trgm`, `unaccent` | +| time-series | `timescaledb` | +| utility | `uuid-ossp` | +| validation | `pg_jsonschema` | + ## Notes - Generated SQL is formatted and readable diff --git a/pkg/writers/pgsql/extensions_test.go b/pkg/writers/pgsql/extensions_test.go new file mode 100644 index 0000000..cfdcc25 --- /dev/null +++ b/pkg/writers/pgsql/extensions_test.go @@ -0,0 +1,260 @@ +package pgsql + +import ( + "reflect" + "strings" + "testing" + + "git.warky.dev/wdevs/relspecgo/pkg/models" + "git.warky.dev/wdevs/relspecgo/pkg/writers" +) + +// buildExtensionSchema returns a single-table schema the extension detection tests mutate. +func buildExtensionSchema(t *testing.T) (*models.Schema, *models.Table) { + t.Helper() + + schema := models.InitSchema("public") + table := models.InitTable("documents", "public") + schema.Tables = append(schema.Tables, table) + return schema, table +} + +func addColumn(table *models.Table, name, sqlType string) *models.Column { + col := models.InitColumn(name, table.Name, table.Schema) + col.Type = sqlType + table.Columns[name] = col + return col +} + +func TestRequiredExtensions_Detection(t *testing.T) { + tests := []struct { + name string + build func(schema *models.Schema, table *models.Table) + want []string + }{ + { + name: "no extensions", + build: func(_ *models.Schema, table *models.Table) { addColumn(table, "id", "integer") }, + want: nil, + }, + { + name: "column type", + build: func(_ *models.Schema, table *models.Table) { + addColumn(table, "embedding", "vector(1536)") + addColumn(table, "name", "citext") + }, + want: []string{"citext", "vector"}, + }, + { + name: "column default function", + build: func(_ *models.Schema, table *models.Table) { + addColumn(table, "id", "uuid").Default = "uuid_generate_v4()" + }, + want: []string{"uuid-ossp"}, + }, + { + name: "check constraint expression", + build: func(_ *models.Schema, table *models.Table) { + addColumn(table, "geom", "geometry") + table.Constraints["chk_geom"] = &models.Constraint{ + Name: "chk_geom", + Type: models.CheckConstraint, + Expression: "ST_IsValid(geom)", + } + }, + want: []string{"postgis"}, + }, + { + name: "partial index predicate", + build: func(_ *models.Schema, table *models.Table) { + addColumn(table, "title", "text") + table.Indexes["idx_title"] = &models.Index{ + Name: "idx_title", + Type: "btree", + Columns: []string{"title"}, + Where: "similarity(title, 'x') > 0.3", + } + }, + want: []string{"pg_trgm"}, + }, + { + name: "view definition", + build: func(schema *models.Schema, table *models.Table) { + addColumn(table, "title", "text") + schema.Views = append(schema.Views, &models.View{ + Name: "v_documents", + Schema: "public", + Definition: "SELECT unaccent(title) FROM documents", + }) + }, + want: []string{"unaccent"}, + }, + { + name: "index access method", + build: func(_ *models.Schema, table *models.Table) { + addColumn(table, "body", "text") + table.Indexes["idx_body"] = &models.Index{ + Name: "idx_body", + Type: "bm25", + Columns: []string{"body"}, + Comment: "with (key_field='id')", + } + }, + want: []string{"pg_search"}, + }, + { + name: "vchord depends on vector", + build: func(_ *models.Schema, table *models.Table) { + addColumn(table, "embedding", "vector(3)") + table.Indexes["idx_embedding"] = &models.Index{ + Name: "idx_embedding", + Type: "vchordrq", + Columns: []string{"embedding"}, + } + }, + want: []string{"vector", "vchord"}, + }, + { + name: "gin on scalar needs btree_gin", + build: func(_ *models.Schema, table *models.Table) { + addColumn(table, "views", "integer") + table.Indexes["idx_views"] = &models.Index{ + Name: "idx_views", + Type: "gin", + Columns: []string{"views"}, + } + }, + want: []string{"btree_gin"}, + }, + { + name: "gist on scalar needs btree_gist", + build: func(_ *models.Schema, table *models.Table) { + addColumn(table, "views", "integer") + table.Indexes["idx_views"] = &models.Index{ + Name: "idx_views", + Type: "gist", + Columns: []string{"views"}, + } + }, + want: []string{"btree_gist"}, + }, + { + name: "gist on geometry uses postgis operator classes", + build: func(_ *models.Schema, table *models.Table) { + addColumn(table, "location", "geometry(Point,4326)") + table.Indexes["idx_location"] = &models.Index{ + Name: "idx_location", + Type: "gist", + Columns: []string{"location"}, + } + }, + want: []string{"postgis"}, + }, + { + name: "gin on jsonb needs no companion", + build: func(_ *models.Schema, table *models.Table) { + addColumn(table, "payload", "jsonb") + table.Indexes["idx_payload"] = &models.Index{ + Name: "idx_payload", + Type: "gin", + Columns: []string{"payload"}, + } + }, + want: nil, + }, + { + name: "gin on text uses pg_trgm", + build: func(_ *models.Schema, table *models.Table) { + addColumn(table, "title", "text") + table.Indexes["idx_title"] = &models.Index{ + Name: "idx_title", + Type: "gin", + Columns: []string{"title"}, + } + }, + want: []string{"pg_trgm"}, + }, + { + name: "gin on array needs no companion", + build: func(_ *models.Schema, table *models.Table) { + addColumn(table, "tags", "text[]") + table.Indexes["idx_tags"] = &models.Index{ + Name: "idx_tags", + Type: "gin", + Columns: []string{"tags"}, + } + }, + want: nil, + }, + { + name: "declared in metadata as string", + build: func(schema *models.Schema, _ *models.Table) { + schema.Metadata = map[string]any{"extensions": "pg_cron, timescaledb"} + }, + want: []string{"pg_cron", "timescaledb"}, + }, + { + name: "declared in metadata as list", + build: func(schema *models.Schema, _ *models.Table) { + schema.Metadata = map[string]any{"extensions": []any{"postgis_topology", "pg_stat_statements"}} + }, + // postgis is pulled in as a dependency of postgis_topology and emitted first. + want: []string{"pg_stat_statements", "postgis", "postgis_topology"}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + schema, table := buildExtensionSchema(t) + tt.build(schema, table) + + if got := requiredExtensions(schema); !reflect.DeepEqual(got, tt.want) { + t.Errorf("requiredExtensions() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestRequiredExtensions_NilSchema(t *testing.T) { + if got := requiredExtensions(nil); got != nil { + t.Errorf("requiredExtensions(nil) = %v, want nil", got) + } +} + +func TestWriteDatabase_QuotesExtensionNames(t *testing.T) { + db := models.InitDatabase("testdb") + schema, table := buildExtensionSchema(t) + addColumn(table, "id", "uuid").Default = "uuid_generate_v4()" + db.Schemas = append(db.Schemas, schema) + + output := writeDatabaseOutput(t, db) + if !strings.Contains(output, `CREATE EXTENSION IF NOT EXISTS "uuid-ossp";`) { + t.Fatalf("expected quoted extension name, got:\n%s", output) + } +} + +func TestGenerateSchemaStatements_ExtensionDependencyOrder(t *testing.T) { + schema, table := buildExtensionSchema(t) + addColumn(table, "embedding", "vector(3)") + table.Indexes["idx_embedding"] = &models.Index{ + Name: "idx_embedding", + Type: "vchordrq", + Columns: []string{"embedding"}, + } + + writer := NewWriter(&writers.WriterOptions{}) + statements, err := writer.GenerateSchemaStatements(schema) + if err != nil { + t.Fatalf("GenerateSchemaStatements failed: %v", err) + } + + joined := strings.Join(statements, "\n") + vector := strings.Index(joined, "CREATE EXTENSION IF NOT EXISTS vector") + vchord := strings.Index(joined, "CREATE EXTENSION IF NOT EXISTS vchord") + if vector < 0 || vchord < 0 { + t.Fatalf("expected vector and vchord extensions, got:\n%s", joined) + } + if vector > vchord { + t.Fatalf("expected vector to be created before vchord, got:\n%s", joined) + } +} diff --git a/pkg/writers/pgsql/migration_writer.go b/pkg/writers/pgsql/migration_writer.go index c8256d3..4b1a9c1 100644 --- a/pkg/writers/pgsql/migration_writer.go +++ b/pkg/writers/pgsql/migration_writer.go @@ -164,14 +164,14 @@ func (w *MigrationWriter) WriteMigration(model *models.Database, current *models func (w *MigrationWriter) generateSchemaScripts(model *models.Schema, current *models.Schema) ([]MigrationScript, error) { scripts := make([]MigrationScript, 0) - if schemaRequiresPGTrgm(model) { + for _, extension := range requiredExtensions(model) { scripts = append(scripts, MigrationScript{ - ObjectName: "extension.pg_trgm", + ObjectName: "extension." + extension, ObjectType: "create extension", Schema: model.Name, Priority: 80, Sequence: len(scripts), - Body: "CREATE EXTENSION IF NOT EXISTS pg_trgm;", + Body: fmt.Sprintf("CREATE EXTENSION IF NOT EXISTS %s;", pgsql.QuoteExtensionName(extension)), }) } @@ -646,13 +646,14 @@ func (w *MigrationWriter) generateIndexScripts(model *models.Schema, current *mo } sql, err := w.executor.ExecuteCreateIndex(CreateIndexData{ - SchemaName: model.Name, - TableName: modelTable.Name, - IndexName: indexName, - IndexType: indexType, - Columns: strings.Join(columnExprs, ", "), - Unique: modelIndex.Unique, - Concurrent: modelIndex.Concurrent, + SchemaName: model.Name, + TableName: modelTable.Name, + IndexName: indexName, + IndexType: indexType, + Columns: strings.Join(columnExprs, ", "), + Unique: modelIndex.Unique, + Concurrent: modelIndex.Concurrent, + StorageParameters: indexStorageParameters(modelIndex.Comment), }) if err != nil { return nil, err @@ -674,20 +675,31 @@ func (w *MigrationWriter) generateIndexScripts(model *models.Schema, current *mo return scripts, nil } +// buildIndexColumnExpressions renders the column list of an index, appending the operator +// class each column needs for the access method (GIN opclasses, pgvector distance ops, +// explicitly requested PostGIS opclasses). Columns that cannot be resolved on the table are +// emitted verbatim. func buildIndexColumnExpressions(table *models.Table, index *models.Index, indexType string) []string { + return buildIndexColumnExpressionsFiltered(table, index, indexType, false) +} + +// buildIndexColumnExpressionsFiltered is buildIndexColumnExpressions with the option to drop +// columns that do not exist on the table instead of emitting them verbatim. +func buildIndexColumnExpressionsFiltered(table *models.Table, index *models.Index, indexType string, skipUnresolved bool) []string { columnExprs := make([]string, 0, len(index.Columns)) for _, colName := range index.Columns { - colExpr := colName - if table != nil { - if col, ok := resolveIndexColumn(table, colName); ok && col != nil { - colExpr = col.SQLName() - if strings.EqualFold(indexType, "gin") { - opClass := ginOperatorClassForColumn(col, index.Comment) - if opClass != "" { - colExpr = fmt.Sprintf("%s %s", col.SQLName(), opClass) - } - } + col, ok := resolveIndexColumn(table, colName) + if !ok || col == nil { + if skipUnresolved { + continue } + columnExprs = append(columnExprs, colName) + continue + } + + colExpr := col.SQLName() + if opClass := indexOperatorClassForColumn(col, indexType, index.Comment); opClass != "" { + colExpr = fmt.Sprintf("%s %s", colExpr, opClass) } columnExprs = append(columnExprs, colExpr) } @@ -1046,5 +1058,19 @@ func indexesEqual(idx1, idx2 *models.Index) bool { return false } } - return true + // Operator class and storage parameters ride along in the index comment. They only + // signal a difference when both sides specify one, so an index whose model side omits + // the hint is not recreated on every migration. + if !indexHintsEqual(extractOperatorClass(idx1.Comment), extractOperatorClass(idx2.Comment)) { + return false + } + return indexHintsEqual(indexStorageParameters(idx1.Comment), indexStorageParameters(idx2.Comment)) +} + +// indexHintsEqual compares two optional index hints, treating an unspecified hint as a match. +func indexHintsEqual(hint1, hint2 string) bool { + if hint1 == "" || hint2 == "" { + return true + } + return strings.EqualFold(hint1, hint2) } diff --git a/pkg/writers/pgsql/migration_writer_test.go b/pkg/writers/pgsql/migration_writer_test.go index d88cb2c..1f5d014 100644 --- a/pkg/writers/pgsql/migration_writer_test.go +++ b/pkg/writers/pgsql/migration_writer_test.go @@ -852,3 +852,93 @@ func TestWriteMigration_NilCurrentTreatsDatabaseAsEmpty(t *testing.T) { t.Fatalf("expected CREATE TABLE in migration output, got:\n%s", output) } } + +func TestWriteMigration_VectorAndPostGISIndexes(t *testing.T) { + current := models.InitDatabase("testdb") + current.Schemas = append(current.Schemas, models.InitSchema("public")) + + model := models.InitDatabase("testdb") + modelSchema := models.InitSchema("public") + + table := models.InitTable("documents", "public") + + embedding := models.InitColumn("embedding", "documents", "public") + embedding.Type = "vector(1536)" + table.Columns["embedding"] = embedding + + location := models.InitColumn("location", "documents", "public") + location.Type = "geometry(Point,4326)" + table.Columns["location"] = location + + table.Indexes["idx_documents_embedding"] = &models.Index{ + Name: "idx_documents_embedding", + Type: "ivfflat", + Columns: []string{"embedding"}, + Comment: "opclass=vector_cosine_ops; with (lists=100)", + } + table.Indexes["idx_documents_location"] = &models.Index{ + Name: "idx_documents_location", + Type: "gist", + Columns: []string{"location"}, + } + + modelSchema.Tables = append(modelSchema.Tables, table) + model.Schemas = append(model.Schemas, modelSchema) + + var buf bytes.Buffer + writer, err := NewMigrationWriter(&writers.WriterOptions{}) + if err != nil { + t.Fatalf("Failed to create writer: %v", err) + } + writer.writer = &buf + + if err := writer.WriteMigration(model, current); err != nil { + t.Fatalf("WriteMigration failed: %v", err) + } + + output := buf.String() + for _, want := range []string{ + "CREATE EXTENSION IF NOT EXISTS postgis;", + "CREATE EXTENSION IF NOT EXISTS vector;", + "vector(1536)", + "geometry(Point,4326)", + "USING ivfflat (embedding vector_cosine_ops) WITH (lists = 100)", + "USING gist (location)", + } { + if !strings.Contains(output, want) { + t.Fatalf("expected migration to contain %q, got:\n%s", want, output) + } + } +} + +func TestIndexesEqual_OperatorClassAndStorageParameters(t *testing.T) { + newIndex := func(comment string) *models.Index { + return &models.Index{ + Name: "idx_documents_embedding", + Type: "hnsw", + Columns: []string{"embedding"}, + Comment: comment, + } + } + + tests := []struct { + name string + comment1 string + comment2 string + wantEqual bool + }{ + {"identical hints", "opclass=vector_l2_ops", "opclass=vector_l2_ops", true}, + {"different operator class", "opclass=vector_l2_ops", "opclass=vector_cosine_ops", false}, + {"different storage parameters", "with (m=16)", "with (m=32)", false}, + {"unspecified hint on one side", "", "opclass=vector_l2_ops; with (m=16)", true}, + {"unrelated comments", "primary lookup index", "primary lookup index", true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := indexesEqual(newIndex(tt.comment1), newIndex(tt.comment2)); got != tt.wantEqual { + t.Errorf("indexesEqual() = %v, want %v", got, tt.wantEqual) + } + }) + } +} diff --git a/pkg/writers/pgsql/templates.go b/pkg/writers/pgsql/templates.go index 16ec99f..7fa738e 100644 --- a/pkg/writers/pgsql/templates.go +++ b/pkg/writers/pgsql/templates.go @@ -140,6 +140,9 @@ type CreateIndexData struct { Columns string Unique bool Concurrent bool + // StorageParameters holds access-method parameters rendered as WITH (...), + // e.g. "lists = 100" for ivfflat or "m = 16, ef_construction = 64" for hnsw. + StorageParameters string } // CreateForeignKeyData contains data for create foreign key template diff --git a/pkg/writers/pgsql/templates/create_index.tmpl b/pkg/writers/pgsql/templates/create_index.tmpl index dda3b01..af57844 100644 --- a/pkg/writers/pgsql/templates/create_index.tmpl +++ b/pkg/writers/pgsql/templates/create_index.tmpl @@ -1,2 +1,2 @@ CREATE {{if .Unique}}UNIQUE {{end}}INDEX {{if .Concurrent}}CONCURRENTLY {{end}}IF NOT EXISTS {{quote_ident .IndexName}} - ON {{qual_table .SchemaName .TableName}} USING {{.IndexType}} ({{.Columns}}); \ No newline at end of file + ON {{qual_table .SchemaName .TableName}} USING {{.IndexType}} ({{.Columns}}){{if .StorageParameters}} WITH ({{.StorageParameters}}){{end}}; \ No newline at end of file diff --git a/pkg/writers/pgsql/writer.go b/pkg/writers/pgsql/writer.go index 969a411..fadf276 100644 --- a/pkg/writers/pgsql/writer.go +++ b/pkg/writers/pgsql/writer.go @@ -6,8 +6,10 @@ import ( "fmt" "io" "os" + "regexp" "sort" "strings" + "sync" "time" "git.warky.dev/wdevs/relspecgo/pkg/models" @@ -147,8 +149,8 @@ func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, erro statements = append(statements, fmt.Sprintf("CREATE SCHEMA IF NOT EXISTS %s", schema.SQLName())) } - if schemaRequiresPGTrgm(schema) { - statements = append(statements, `CREATE EXTENSION IF NOT EXISTS pg_trgm`) + for _, extension := range requiredExtensions(schema) { + statements = append(statements, fmt.Sprintf("CREATE EXTENSION IF NOT EXISTS %s", pgsql.QuoteExtensionName(extension))) } // Phase 2: Create sequences @@ -271,18 +273,12 @@ func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, erro indexType = "btree" } - // Build column expressions with operator class support for GIN indexes - columnExprs := make([]string, 0, len(index.Columns)) - for _, colName := range index.Columns { - colExpr := colName - if col, ok := resolveIndexColumn(table, colName); ok { - if strings.EqualFold(indexType, "gin") { - if opClass := ginOperatorClassForColumn(col, index.Comment); opClass != "" { - colExpr = fmt.Sprintf("%s %s", colName, opClass) - } - } - } - columnExprs = append(columnExprs, colExpr) + // Build column expressions with operator class support (GIN, pgvector, PostGIS) + columnExprs := buildIndexColumnExpressions(table, index, indexType) + + withClause := "" + if params := indexStorageParameters(index.Comment); params != "" { + withClause = fmt.Sprintf(" WITH (%s)", params) } whereClause := "" @@ -290,8 +286,8 @@ func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, erro whereClause = fmt.Sprintf(" WHERE %s", index.Where) } - stmt := fmt.Sprintf("CREATE %sINDEX IF NOT EXISTS %s ON %s USING %s (%s)%s", - uniqueStr, quoteIdentifier(index.Name), w.qualTable(schema.SQLName(), table.SQLName()), indexType, strings.Join(columnExprs, ", "), whereClause) + stmt := fmt.Sprintf("CREATE %sINDEX IF NOT EXISTS %s ON %s USING %s (%s)%s%s", + uniqueStr, quoteIdentifier(index.Name), w.qualTable(schema.SQLName(), table.SQLName()), indexType, strings.Join(columnExprs, ", "), withClause, whereClause) statements = append(statements, stmt) } } @@ -819,11 +815,14 @@ func (w *Writer) writeCreateSchema(schema *models.Schema) error { } func (w *Writer) writeRequiredExtensions(schema *models.Schema) error { - if !schemaRequiresPGTrgm(schema) { + extensions := requiredExtensions(schema) + if len(extensions) == 0 { return nil } - fmt.Fprintln(w.writer, "CREATE EXTENSION IF NOT EXISTS pg_trgm;") + for _, extension := range extensions { + fmt.Fprintf(w.writer, "CREATE EXTENSION IF NOT EXISTS %s;\n", pgsql.QuoteExtensionName(extension)) + } fmt.Fprintln(w.writer) return nil } @@ -1063,21 +1062,13 @@ func (w *Writer) writeIndexes(schema *models.Schema) error { indexName = fmt.Sprintf("%s_%s_%s", indexType, table.SQLName(), strings.ToLower(columnSuffix)) } - // Build column list with operator class support for GIN indexes - columnExprs := make([]string, 0, len(index.Columns)) - for _, colName := range index.Columns { - if col, ok := resolveIndexColumn(table, colName); ok { - colExpr := col.SQLName() - if strings.EqualFold(index.Type, "gin") { - opClass := ginOperatorClassForColumn(col, index.Comment) - if opClass != "" { - colExpr = fmt.Sprintf("%s %s", col.SQLName(), opClass) - } - } - columnExprs = append(columnExprs, colExpr) - } + indexType := index.Type + if indexType == "" { + indexType = "btree" } + // Build column list with operator class support (GIN, pgvector, PostGIS) + columnExprs := buildIndexColumnExpressionsFiltered(table, index, indexType, true) if len(columnExprs) == 0 { continue } @@ -1087,9 +1078,9 @@ func (w *Writer) writeIndexes(schema *models.Schema) error { unique = "UNIQUE " } - indexType := index.Type - if indexType == "" { - indexType = "btree" + withClause := "" + if params := indexStorageParameters(index.Comment); params != "" { + withClause = fmt.Sprintf(" WITH (%s)", params) } whereClause := "" @@ -1104,8 +1095,8 @@ func (w *Writer) writeIndexes(schema *models.Schema) error { fmt.Fprintf(w.writer, "CREATE %sINDEX %sIF NOT EXISTS %s\n", unique, concurrently, indexName) - fmt.Fprintf(w.writer, " ON %s USING %s (%s)%s;\n\n", - w.qualTable(schema.SQLName(), table.SQLName()), indexType, strings.Join(columnExprs, ", "), whereClause) + fmt.Fprintf(w.writer, " ON %s USING %s (%s)%s%s;\n\n", + w.qualTable(schema.SQLName(), table.SQLName()), indexType, strings.Join(columnExprs, ", "), withClause, whereClause) } } @@ -1483,7 +1474,69 @@ func isTextTypeWithoutLength(colType string) bool { return strings.EqualFold(colType, "text") } -func ginOperatorClassForColumn(col *models.Column, comment string) string { +// vectorOperatorClasses maps pgvector operator classes to the column base type they +// apply to. pgvector defines no default operator class, so an hnsw/ivfflat index must +// always name one explicitly. +var vectorOperatorClasses = map[string]string{ + "vector_l2_ops": "vector", + "vector_ip_ops": "vector", + "vector_cosine_ops": "vector", + "vector_l1_ops": "vector", + "halfvec_l2_ops": "halfvec", + "halfvec_ip_ops": "halfvec", + "halfvec_cosine_ops": "halfvec", + "halfvec_l1_ops": "halfvec", + "sparsevec_l2_ops": "sparsevec", + "sparsevec_ip_ops": "sparsevec", + "sparsevec_cosine_ops": "sparsevec", + "sparsevec_l1_ops": "sparsevec", + "bit_hamming_ops": "bit", + "bit_jaccard_ops": "bit", +} + +// defaultVectorOperatorClasses is the operator class used for an hnsw/ivfflat index when +// the index comment does not request one. Cosine distance is the common default for +// embedding columns; override it with an "opclass" hint in the index comment. +var defaultVectorOperatorClasses = map[string]string{ + "vector": "vector_cosine_ops", + "halfvec": "halfvec_cosine_ops", + "sparsevec": "sparsevec_cosine_ops", + "bit": "bit_hamming_ops", +} + +// spatialOperatorClasses are the PostGIS operator classes recognized in index comments. +// PostGIS installs default operator classes for gist/spgist/brin, so these are only +// emitted when explicitly requested (e.g. the 3D/nD variants). +var spatialOperatorClasses = map[string]bool{ + "gist_geometry_ops_2d": true, + "gist_geometry_ops_nd": true, + "gist_geography_ops": true, + "spgist_geometry_ops_2d": true, + "spgist_geometry_ops_3d": true, + "spgist_geometry_ops_nd": true, + "brin_geometry_inclusion_ops_2d": true, + "brin_geometry_inclusion_ops_3d": true, + "brin_geometry_inclusion_ops_4d": true, + "brin_geography_inclusion_ops_2d": true, + "btree_geometry_ops": true, + "btree_geography_ops": true, +} + +// isVectorIndexMethod reports whether the access method indexes pgvector types, which +// covers both pgvector itself (hnsw, ivfflat) and VectorChord (vchordrq, vchordg). +func isVectorIndexMethod(method string) bool { + switch strings.ToLower(strings.TrimSpace(method)) { + case "hnsw", "ivfflat", "vchordrq", "vchordg": + return true + default: + return false + } +} + +// indexOperatorClassForColumn returns the operator class to emit for a column in an index +// of the given access method, honouring an explicit request from the index comment when it +// is compatible with the column type. +func indexOperatorClassForColumn(col *models.Column, indexType, comment string) string { if col == nil { return "" } @@ -1492,26 +1545,53 @@ func ginOperatorClassForColumn(col *models.Column, comment string) string { baseType := pgsql.CanonicalizeBaseType(pgsql.ExtractBaseTypeLower(sqlType)) isArray := pgsql.IsArrayType(sqlType) requested := extractOperatorClass(comment) - - if requested != "" && ginOperatorClassCompatible(baseType, isArray, requested) { - return requested + method := strings.ToLower(strings.TrimSpace(indexType)) + if method == "" { + method = "btree" } - if isArray { - return "array_ops" + if requested != "" && operatorClassCompatible(method, baseType, isArray, requested) { + return requested } switch { - case isTextGinBaseType(baseType): - return "gin_trgm_ops" - case baseType == "jsonb": - return "jsonb_ops" + case method == "gin": + if isArray { + return "array_ops" + } + switch { + case isTextGinBaseType(baseType): + return "gin_trgm_ops" + case baseType == "jsonb": + return "jsonb_ops" + default: + return requested + } + case isVectorIndexMethod(method): + if isArray { + return "" + } + return defaultVectorOperatorClasses[baseType] default: - return requested + // gist/spgist/brin/btree have default operator classes (PostGIS included), + // so nothing is emitted unless the comment requested a compatible class. + return "" } } -func ginOperatorClassCompatible(baseType string, isArray bool, opClass string) bool { +// ginOperatorClassForColumn is the GIN-specific form of indexOperatorClassForColumn. +func ginOperatorClassForColumn(col *models.Column, comment string) string { + return indexOperatorClassForColumn(col, "gin", comment) +} + +func operatorClassCompatible(method, baseType string, isArray bool, opClass string) bool { + if vectorType, ok := vectorOperatorClasses[opClass]; ok { + return !isArray && baseType == vectorType && isVectorIndexMethod(method) + } + if spatialOperatorClasses[opClass] { + return !isArray && pgsql.IsSpatialType(baseType) + } + switch opClass { case "gin_trgm_ops", "gin_bigm_ops": return !isArray && isTextGinBaseType(baseType) @@ -1524,6 +1604,10 @@ func ginOperatorClassCompatible(baseType string, isArray bool, opClass string) b } } +func ginOperatorClassCompatible(baseType string, isArray bool, opClass string) bool { + return operatorClassCompatible("gin", baseType, isArray, opClass) +} + func isTextGinBaseType(baseType string) bool { switch baseType { case "text", "varchar", "character varying", "char", "character", "string", "citext", "bpchar": @@ -1533,29 +1617,188 @@ func isTextGinBaseType(baseType string) bool { } } -func schemaRequiresPGTrgm(schema *models.Schema) bool { +// requiredExtensions returns the PostgreSQL extensions a schema depends on, ordered so +// that dependencies are created first (postgis before postgis_topology, vector before +// vchord). Extensions are detected from column types, index access methods, resolved +// operator classes, and function calls in defaults, check constraints, partial index +// predicates and view definitions. Extensions that leave no trace in the model (pg_cron, +// timescaledb, postgres_fdw, …) can be declared in schema.Metadata["extensions"]. +func requiredExtensions(schema *models.Schema) []string { if schema == nil { - return false + return nil } + + required := make(map[string]bool) + add := func(names ...string) { + for _, name := range names { + if name != "" { + required[name] = true + } + } + } + + add(declaredExtensions(schema)...) + + for _, view := range schema.Views { + if view == nil { + continue + } + add(pgsql.ExtensionsForExpression(view.Definition)...) + } + for _, table := range schema.Tables { if table == nil { continue } - for _, index := range table.Indexes { - if index == nil || !strings.EqualFold(index.Type, "gin") { + + for _, col := range table.Columns { + if col == nil { continue } + add(pgsql.TypeExtension(effectiveColumnSQLType(col))) + if def, ok := col.Default.(string); ok { + add(pgsql.ExtensionsForExpression(def)...) + } + } + + for _, constraint := range table.Constraints { + if constraint == nil { + continue + } + add(pgsql.ExtensionsForExpression(constraint.Expression)...) + } + + for _, index := range table.Indexes { + if index == nil { + continue + } + add(pgsql.IndexMethodExtension(index.Type)) + add(pgsql.ExtensionsForExpression(index.Where)...) + for _, colName := range index.Columns { col, ok := resolveIndexColumn(table, colName) if !ok || col == nil { continue } - if ginOperatorClassForColumn(col, index.Comment) == "gin_trgm_ops" { - return true - } + opClass := indexOperatorClassForColumn(col, index.Type, index.Comment) + add(pgsql.OperatorClassExtension(opClass)) + add(btreeCompanionExtension(index.Type, col, opClass)) } } } + + extensions := make([]string, 0, len(required)) + for ext := range required { + extensions = append(extensions, ext) + } + + // Pull in dependencies, so a declared postgis_topology also creates postgis. + for i := 0; i < len(extensions); i++ { + for _, dependency := range pgsql.ExtensionDependencies(extensions[i]) { + if !required[dependency] { + required[dependency] = true + extensions = append(extensions, dependency) + } + } + } + + return pgsql.SortExtensions(extensions) +} + +// declaredExtensions reads schema.Metadata["extensions"], which accepts either a list or a +// comma-separated string. Unknown names are kept: the metadata is an explicit instruction. +func declaredExtensions(schema *models.Schema) []string { + value, ok := schema.Metadata["extensions"] + if !ok { + return nil + } + + var names []string + switch declared := value.(type) { + case string: + names = strings.Split(declared, ",") + case []string: + names = declared + case []any: + for _, item := range declared { + if name, ok := item.(string); ok { + names = append(names, name) + } + } + default: + return nil + } + + cleaned := make([]string, 0, len(names)) + for _, name := range names { + if name = strings.TrimSpace(name); name != "" { + cleaned = append(cleaned, name) + } + } + return cleaned +} + +// btreeCompanionExtension returns btree_gin or btree_gist when a GIN/GiST index covers a +// scalar type that neither access method has a built-in operator class for. Without the +// companion extension PostgreSQL rejects the CREATE INDEX outright. +func btreeCompanionExtension(indexType string, col *models.Column, opClass string) string { + if opClass != "" { + return "" + } + + method := strings.ToLower(strings.TrimSpace(indexType)) + if method != "gin" && method != "gist" { + return "" + } + + sqlType := effectiveColumnSQLType(col) + if pgsql.IsArrayType(sqlType) { + return "" + } + + baseType := pgsql.CanonicalizeBaseType(pgsql.ExtractBaseTypeLower(sqlType)) + if pgsql.TypeExtension(baseType) != "" { + // Extension types (geometry, vector, citext, …) ship their own operator classes. + return "" + } + + if method == "gin" { + if nativeGinBaseType(baseType) { + return "" + } + return "btree_gin" + } + if nativeGistBaseType(baseType) { + return "" + } + return "btree_gist" +} + +// nativeGinBaseType reports whether core PostgreSQL provides a GIN operator class. +func nativeGinBaseType(baseType string) bool { + switch baseType { + case "jsonb", "json", "tsvector", "tsquery": + return true + default: + return false + } +} + +// nativeGistBaseType reports whether core PostgreSQL provides a GiST operator class. +func nativeGistBaseType(baseType string) bool { + switch baseType { + case "tsvector", "tsquery", "point", "box", "circle", "polygon", "line", "lseg", "path", "inet", "cidr": + return true + } + return strings.HasSuffix(baseType, "range") || strings.HasSuffix(baseType, "multirange") +} + +func schemaRequiresPGTrgm(schema *models.Schema) bool { + for _, ext := range requiredExtensions(schema) { + if ext == "pg_trgm" { + return true + } + } return false } @@ -1642,14 +1885,21 @@ func formatStringList(items []string) string { // extractOperatorClass extracts operator class from index comment/note // Looks for common operator classes like gin_trgm_ops, gist_trgm_ops, etc. +// explicitOperatorClassPattern matches an "opclass=" hint, the form the PostgreSQL +// reader uses to carry an index's operator class through the model. +var explicitOperatorClassPattern = regexp.MustCompile(`(?i)\bopclass\s*=\s*([a-z_][a-z0-9_]*)\b`) + func extractOperatorClass(comment string) string { if comment == "" { return "" } + lowerComment := strings.ToLower(comment) - // Common GIN/GiST operator classes - opClasses := []string{"gin_trgm_ops", "gist_trgm_ops", "gin_bigm_ops", "jsonb_ops", "jsonb_path_ops", "array_ops"} - for _, op := range opClasses { + if matches := explicitOperatorClassPattern.FindStringSubmatch(lowerComment); len(matches) > 1 { + return matches[1] + } + + for _, op := range knownOperatorClasses() { if strings.Contains(lowerComment, op) { return op } @@ -1657,6 +1907,35 @@ func extractOperatorClass(comment string) string { return "" } +// knownOperatorClasses lists every operator class recognized in an index comment, +// longest name first so that e.g. gist_geometry_ops_nd wins over a shorter prefix. +var knownOperatorClasses = sync.OnceValue(func() []string { + names := []string{"gin_trgm_ops", "gist_trgm_ops", "gin_bigm_ops", "jsonb_ops", "jsonb_path_ops", "array_ops"} + for name := range vectorOperatorClasses { + names = append(names, name) + } + for name := range spatialOperatorClasses { + names = append(names, name) + } + sort.Slice(names, func(i, j int) bool { + if len(names[i]) != len(names[j]) { + return len(names[i]) > len(names[j]) + } + return names[i] < names[j] + }) + return names +}) + +// indexStorageParameters extracts access-method storage parameters from an index comment. +// Only well-formed "key = value" pairs are kept, so comment prose cannot leak into DDL. +// Example: "opclass=vector_cosine_ops with (m=16, ef_construction=64)" -> "m = 16, ef_construction = 64". +func indexStorageParameters(comment string) string { + if comment == "" { + return "" + } + return pgsql.FormatStorageParameters(pgsql.ExtractWithClause(comment)) +} + // escapeQuote escapes single quotes in strings for SQL func escapeQuote(s string) string { return strings.ReplaceAll(s, "'", "''") diff --git a/pkg/writers/pgsql/writer_test.go b/pkg/writers/pgsql/writer_test.go index d217ca7..fdab68a 100644 --- a/pkg/writers/pgsql/writer_test.go +++ b/pkg/writers/pgsql/writer_test.go @@ -1310,3 +1310,203 @@ func TestWriteSchema_UsesStorageTypeForSerialAlterStatements(t *testing.T) { t.Fatalf("expected serial alter to include USING cast, got:\n%s", output) } } + +// buildVectorSpatialSchema returns a database with a pgvector column and a PostGIS column. +func buildVectorSpatialSchema(indexType, indexComment string) *models.Database { + db := models.InitDatabase("testdb") + schema := models.InitSchema("public") + + table := models.InitTable("documents", "public") + + embedding := models.InitColumn("embedding", "documents", "public") + embedding.Type = "vector(1536)" + table.Columns["embedding"] = embedding + + location := models.InitColumn("location", "documents", "public") + location.Type = "geometry(Point,4326)" + table.Columns["location"] = location + + if indexType != "" { + index := &models.Index{ + Name: "idx_documents_embedding", + Type: indexType, + Columns: []string{"embedding"}, + Comment: indexComment, + } + table.Indexes[index.Name] = index + } + + schema.Tables = append(schema.Tables, table) + db.Schemas = append(db.Schemas, schema) + return db +} + +func writeDatabaseOutput(t *testing.T, db *models.Database) string { + t.Helper() + + var buf bytes.Buffer + writer := NewWriter(&writers.WriterOptions{}) + writer.writer = &buf + + if err := writer.WriteDatabase(db); err != nil { + t.Fatalf("WriteDatabase failed: %v", err) + } + return buf.String() +} + +func TestWriteDatabase_VectorAndPostGISColumnsCreateExtensions(t *testing.T) { + output := writeDatabaseOutput(t, buildVectorSpatialSchema("", "")) + + for _, want := range []string{ + "CREATE EXTENSION IF NOT EXISTS postgis;", + "CREATE EXTENSION IF NOT EXISTS vector;", + "vector(1536)", + "geometry(Point,4326)", + } { + if !strings.Contains(output, want) { + t.Fatalf("expected output to contain %q, got:\n%s", want, output) + } + } + + // postgis must be created before postgis-dependent extensions and stay deterministic + if strings.Index(output, "EXISTS postgis;") > strings.Index(output, "EXISTS vector;") { + t.Fatalf("expected extensions to be emitted in sorted order, got:\n%s", output) + } +} + +func TestWriteDatabase_HNSWIndexUsesDefaultVectorOperatorClass(t *testing.T) { + output := writeDatabaseOutput(t, buildVectorSpatialSchema("hnsw", "")) + + if !strings.Contains(output, "USING hnsw (embedding vector_cosine_ops)") { + t.Fatalf("expected hnsw index with default vector operator class, got:\n%s", output) + } + if !strings.Contains(output, "CREATE EXTENSION IF NOT EXISTS vector;") { + t.Fatalf("expected pgvector extension, got:\n%s", output) + } +} + +func TestWriteDatabase_VectorIndexHonoursRequestedOperatorClassAndStorageParameters(t *testing.T) { + output := writeDatabaseOutput(t, buildVectorSpatialSchema("ivfflat", "opclass=vector_l2_ops; with (lists=100)")) + + if !strings.Contains(output, "USING ivfflat (embedding vector_l2_ops) WITH (lists = 100)") { + t.Fatalf("expected ivfflat index with requested opclass and storage parameters, got:\n%s", output) + } +} + +func TestWriteDatabase_VectorIndexIgnoresIncompatibleOperatorClass(t *testing.T) { + output := writeDatabaseOutput(t, buildVectorSpatialSchema("hnsw", "opclass=halfvec_l2_ops")) + + if !strings.Contains(output, "USING hnsw (embedding vector_cosine_ops)") { + t.Fatalf("expected halfvec operator class to be rejected for a vector column, got:\n%s", output) + } +} + +func TestWriteDatabase_VectorIndexIgnoresCommentProseInStorageParameters(t *testing.T) { + output := writeDatabaseOutput(t, buildVectorSpatialSchema("hnsw", "tuned with (m=16, ef_construction=64, drop table foo)")) + + if !strings.Contains(output, "WITH (m = 16, ef_construction = 64)") { + t.Fatalf("expected only well-formed storage parameters, got:\n%s", output) + } + if strings.Contains(output, "drop table") { + t.Fatalf("expected prose to be dropped from storage parameters, got:\n%s", output) + } +} + +func TestWriteDatabase_GistIndexOnGeometryUsesDefaultOperatorClass(t *testing.T) { + db := models.InitDatabase("testdb") + schema := models.InitSchema("public") + + table := models.InitTable("places", "public") + geom := models.InitColumn("geom", "places", "public") + geom.Type = "geometry(Point,4326)" + table.Columns["geom"] = geom + + table.Indexes["idx_places_geom"] = &models.Index{ + Name: "idx_places_geom", + Type: "gist", + Columns: []string{"geom"}, + } + + schema.Tables = append(schema.Tables, table) + db.Schemas = append(db.Schemas, schema) + + output := writeDatabaseOutput(t, db) + + if !strings.Contains(output, "USING gist (geom)") { + t.Fatalf("expected gist index to rely on the PostGIS default operator class, got:\n%s", output) + } +} + +func TestWriteDatabase_GistIndexHonoursRequestedSpatialOperatorClass(t *testing.T) { + db := models.InitDatabase("testdb") + schema := models.InitSchema("public") + + table := models.InitTable("places", "public") + geom := models.InitColumn("geom", "places", "public") + geom.Type = "geometry(PointZ,4326)" + table.Columns["geom"] = geom + + table.Indexes["idx_places_geom_nd"] = &models.Index{ + Name: "idx_places_geom_nd", + Type: "gist", + Columns: []string{"geom"}, + Comment: "opclass=gist_geometry_ops_nd", + } + + schema.Tables = append(schema.Tables, table) + db.Schemas = append(db.Schemas, schema) + + output := writeDatabaseOutput(t, db) + + if !strings.Contains(output, "USING gist (geom gist_geometry_ops_nd)") { + t.Fatalf("expected requested spatial operator class, got:\n%s", output) + } +} + +func TestGenerateDatabaseStatements_VectorIndexIncludesOperatorClassAndParameters(t *testing.T) { + db := buildVectorSpatialSchema("hnsw", "opclass=vector_ip_ops; with (m=16)") + + writer := NewWriter(&writers.WriterOptions{}) + statements, err := writer.GenerateDatabaseStatements(db) + if err != nil { + t.Fatalf("GenerateDatabaseStatements failed: %v", err) + } + + joined := strings.Join(statements, "\n") + for _, want := range []string{ + "CREATE EXTENSION IF NOT EXISTS vector", + "CREATE EXTENSION IF NOT EXISTS postgis", + "USING hnsw (embedding vector_ip_ops) WITH (m = 16)", + } { + if !strings.Contains(joined, want) { + t.Fatalf("expected statements to contain %q, got:\n%s", want, joined) + } + } +} + +func TestIndexStorageParameters(t *testing.T) { + tests := []struct { + name string + comment string + want string + }{ + {"empty", "", ""}, + {"no with clause", "opclass=vector_cosine_ops", ""}, + {"single parameter", "with (lists=100)", "lists = 100"}, + {"multiple parameters", "WITH (m = 16, ef_construction = 64)", "m = 16, ef_construction = 64"}, + {"quoted value kept", "with (fillfactor='90')", "fillfactor = '90'"}, + {"bm25 key field", "with (key_field='id')", "key_field = 'id'"}, + {"dollar quoted value", "with (options = $$[build.internal]\nlists = [4096]$$)", "options = $$[build.internal]\nlists = [4096]$$"}, + {"dollar quoted value with parens", "with (options = $$f(x)$$, m = 16)", "options = $$f(x)$$, m = 16"}, + {"prose dropped", "with (lists=100, please drop everything)", "lists = 100"}, + {"unterminated quote dropped", "with (key_field='id)", ""}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := indexStorageParameters(tt.comment); got != tt.want { + t.Errorf("indexStorageParameters(%q) = %q, want %q", tt.comment, got, tt.want) + } + }) + } +}