From f259df125899cff67138f3024e7bf6a372f19a7f Mon Sep 17 00:00:00 2001 From: Hein Date: Sat, 29 Aug 2026 21:27:42 +0200 Subject: [PATCH] feat(spectypes): add support for PostGIS and pgvector types * Implement custom types: SqlGeometry, SqlGeography, SqlHalfVector, SqlSparseVector, SqlBitVector * Add spatial filter operators and vector similarity operators * Include metadata and OpenAPI reporting for geometry/vector column types * Create tests for EWKB and WKT conversions --- README.md | 67 +++ pkg/common/spatial_helpers.go | 317 +++++++++++ pkg/common/spatial_helpers_test.go | 145 +++++ pkg/common/types.go | 15 + pkg/openapi/generator.go | 23 + pkg/reflection/model_utils_geo_test.go | 56 ++ pkg/reflection/spectypes_helpers.go | 94 ++++ pkg/resolvespec/filter_test.go | 37 ++ pkg/resolvespec/handler.go | 78 ++- pkg/restheadspec/handler.go | 56 +- pkg/restheadspec/headers.go | 101 ++++ pkg/restheadspec/spatial_filter_test.go | 124 +++++ pkg/spectypes/sql_geo_types.go | 239 ++++++++ pkg/spectypes/sql_geo_types_test.go | 126 +++++ pkg/spectypes/sql_types_test.go | 1 - pkg/spectypes/sql_vector_types.go | 308 +++++++++++ pkg/spectypes/sql_vector_types_test.go | 149 +++++ pkg/spectypes/type_names.go | 93 ++++ pkg/spectypes/type_names_test.go | 66 +++ pkg/spectypes/wkb.go | 697 ++++++++++++++++++++++++ pkg/spectypes/wkb_test.go | 91 ++++ todo.md | 12 +- 22 files changed, 2886 insertions(+), 9 deletions(-) create mode 100644 pkg/common/spatial_helpers.go create mode 100644 pkg/common/spatial_helpers_test.go create mode 100644 pkg/reflection/model_utils_geo_test.go create mode 100644 pkg/reflection/spectypes_helpers.go create mode 100644 pkg/restheadspec/spatial_filter_test.go create mode 100644 pkg/spectypes/sql_geo_types.go create mode 100644 pkg/spectypes/sql_geo_types_test.go create mode 100644 pkg/spectypes/sql_vector_types.go create mode 100644 pkg/spectypes/sql_vector_types_test.go create mode 100644 pkg/spectypes/type_names.go create mode 100644 pkg/spectypes/type_names_test.go create mode 100644 pkg/spectypes/wkb.go create mode 100644 pkg/spectypes/wkb_test.go diff --git a/README.md b/README.md index 3fbb8bf..4354c19 100644 --- a/README.md +++ b/README.md @@ -135,6 +135,73 @@ For complete documentation including setup, headers, lifecycle hooks, cursor pag For detailed examples of reading data, cursor pagination, recursive CRUD operations, filtering, sorting, and more, see [pkg/resolvespec/README.md](pkg/resolvespec/README.md). +## PostGIS & Vector (PostgreSQL only) + +First-class support for PostGIS geometry/geography and pgvector columns in `resolvespec` + `restheadspec`. No extra dependencies. On non-Postgres databases the spatial/vector operators simply don't match. + +### Column types (`pkg/spectypes`) + +| Go type | SQL type | Wire / JSON | +|--------------------|--------------|--------------------------------------------------------| +| `SqlGeometry` | `geometry` | JSON in/out = **GeoJSON**; also accepts EWKT / hex-EWKB | +| `SqlGeography` | `geography` | same as `SqlGeometry` | +| `SqlVector` | `vector` | `[]float32` ⇄ `[1,2,3]` | +| `SqlHalfVector` | `halfvec` | `[]float32` ⇄ `[1,2,3]` | +| `SqlSparseVector` | `sparsevec` | `{"dim":8,"indices":[1,4],"values":[0.5,0.2]}` | +| `SqlBitVector` | `bit`/`varbit` | bool array or `"1011"` string | + +- Geometry `Value()` emits `SRID=;` (PostGIS implicit text→geometry cast; no wrapper function needed). +- Declare dimensioned types with a tag: `gorm:"type:vector(1536)"` — the tag wins over the canonical name in metadata/OpenAPI. +- Metadata endpoint and OpenAPI schema report `geometry`/`vector`/`halfvec`/`sparsevec`/`bit`. + +### Spatial filter operators + +`value` is a geometry (GeoJSON object, EWKT string, or hex-EWKB) unless noted. + +| Operator | Value shape | +|----------|-------------| +| `st_intersects`, `st_contains`, `st_within`, `st_covers`, `st_coveredby`, `st_overlaps`, `st_touches`, `st_crosses`, `st_equals`, `st_disjoint` | geometry | +| `st_dwithin` | `{"geom": , "distance": }` | +| `bbox` (alias `&&`) | geometry, or `{"bbox":[minx,miny,maxx,maxy],"srid":4326}` | + +### Vector similarity filter operators + +| Operator | pgvector op | Value shape | +|----------|-------------|-------------| +| `l2_within` / `euclidean_within` | `<->` | `{"vector":[...], "distance": }` | +| `cosine_within` | `<=>` | same (also `"lt"`/`"lte"`/`"gt"`/`"gte"` instead of `"distance"`) | +| `ip_within` / `inner_within` | `<#>` | same | + +### KNN search (ordering + distance column) + +**resolvespec** — `options.vector_search`: + +```json +{ "options": { "vector_search": { + "column": "embedding", + "vector": [0.1, 0.2, 0.3], + "metric": "cosine", + "as": "_distance", + "direction": "asc" +}}} +``` + +Orders rows by distance; when `as` is set, returns the distance as an extra column (all model columns are auto-selected). +`metric`: `l2` (default) | `cosine` | `ip`. + +**restheadspec** — headers: + +```HTTP +X-Vector-Search-embedding: cosine +X-Vector-Search-Vector: [0.1,0.2,0.3] +X-Vector-Search-As: _distance +X-Vector-Search-Dir: asc +``` + +Spatial/vector filters via headers: `X-SpatialFilter-` / `X-VectorFilter-` with a JSON operator object, e.g. +`X-SpatialFilter-geom: {"op":"st_dwithin","geom":"SRID=4326;POINT(0 0)","distance":1000}` +(optional `"logic":"or"`). + ## Installation ```Shell diff --git a/pkg/common/spatial_helpers.go b/pkg/common/spatial_helpers.go new file mode 100644 index 0000000..98eae6a --- /dev/null +++ b/pkg/common/spatial_helpers.go @@ -0,0 +1,317 @@ +package common + +import ( + "encoding/hex" + "encoding/json" + "fmt" + "strconv" + "strings" +) + +// This file implements PostGIS spatial and pgvector similarity filter operators. +// The builders return parameterised SQL fragments (with `?` placeholders) plus +// their args, matching the style of BuildInCondition / BuildArrayOverlapCondition. +// PostgreSQL only — on other databases these operators simply will not resolve. + +// ── vector similarity ─────────────────────────────────────────────────────── + +// VectorOperator maps a metric name to its pgvector distance operator. +// +// "l2" / "euclidean" / "" -> <-> +// "cosine" -> <=> +// "ip" / "inner" / "dot" -> <#> +func VectorOperator(metric string) string { + switch strings.ToLower(strings.TrimSpace(metric)) { + case "cosine", "cos": + return "<=>" + case "ip", "inner", "dot", "innerproduct", "inner_product": + return "<#>" + default: + return "<->" + } +} + +// VectorLiteral converts a vector value into a pgvector literal string +// "[1,2,3]". Accepts []float32, []float64, []int, []any (of numbers), or an +// already-formatted string. +func VectorLiteral(value any) (string, error) { + switch v := value.(type) { + case string: + s := strings.TrimSpace(v) + if strings.HasPrefix(s, "[") && strings.HasSuffix(s, "]") { + return s, nil + } + return "", fmt.Errorf("vector literal: malformed string %q", v) + case []float32: + return floatsToVectorLiteral(len(v), func(i int) float64 { return float64(v[i]) }), nil + case []float64: + return floatsToVectorLiteral(len(v), func(i int) float64 { return v[i] }), nil + case []int: + return floatsToVectorLiteral(len(v), func(i int) float64 { return float64(v[i]) }), nil + case []any: + nums := make([]float64, len(v)) + for i, e := range v { + f, ok := toFloat(e) + if !ok { + return "", fmt.Errorf("vector literal: element %d is not a number (%T)", i, e) + } + nums[i] = f + } + return floatsToVectorLiteral(len(nums), func(i int) float64 { return nums[i] }), nil + default: + return "", fmt.Errorf("vector literal: unsupported type %T", value) + } +} + +func floatsToVectorLiteral(n int, at func(int) float64) string { + parts := make([]string, n) + for i := 0; i < n; i++ { + parts[i] = strconv.FormatFloat(at(i), 'f', -1, 32) + } + return "[" + strings.Join(parts, ",") + "]" +} + +// BuildVectorCondition builds a pgvector distance-threshold filter. +// +// operator: "l2_within" | "cosine_within" | "ip_within" +// value: {"vector": [...], "distance": } +// {"vector": [...], "lt"|"lte"|"gt"|"gte": } +// +// Produces e.g. `embedding <=> ? < ?` with args [vectorLiteral, threshold]. +func BuildVectorCondition(column, operator string, value any) (query string, args []interface{}, ok bool) { + var op string + switch strings.ToLower(operator) { + case "l2_within", "l2distance_within", "euclidean_within": + op = "<->" + case "cosine_within", "cosinedistance_within": + op = "<=>" + case "ip_within", "inner_within", "negativeinnerproduct_within": + op = "<#>" + default: + return "", nil, false + } + + m, mok := value.(map[string]any) + if !mok { + return "", nil, false + } + lit, err := VectorLiteral(m["vector"]) + if err != nil { + return "", nil, false + } + + cmp := "<" + var threshold any + if t, ok := m["distance"]; ok { + threshold = t + } else { + for _, k := range []string{"lt", "lte", "gt", "gte"} { + if t, ok := m[k]; ok { + threshold = t + switch k { + case "lt": + cmp = "<" + case "lte": + cmp = "<=" + case "gt": + cmp = ">" + case "gte": + cmp = ">=" + } + break + } + } + } + if threshold == nil { + return "", nil, false + } + f, fok := toFloat(threshold) + if !fok { + return "", nil, false + } + + return fmt.Sprintf("%s %s ? %s ?", column, op, cmp), []interface{}{lit, f}, true +} + +// ── PostGIS spatial ───────────────────────────────────────────────────────── + +var spatialPredicates = map[string]string{ + "st_intersects": "ST_Intersects", + "st_contains": "ST_Contains", + "st_within": "ST_Within", + "st_covers": "ST_Covers", + "st_coveredby": "ST_CoveredBy", + "st_overlaps": "ST_Overlaps", + "st_touches": "ST_Touches", + "st_crosses": "ST_Crosses", + "st_equals": "ST_Equals", + "st_disjoint": "ST_Disjoint", +} + +// BuildSpatialCondition builds a PostGIS spatial filter. +// +// "st_dwithin" value: {"geom": , "distance": } +// "st_intersects" / "st_contains" / "st_within" / "st_covers" / +// "st_coveredby" / "st_overlaps" / "st_touches" / "st_crosses" / +// "st_equals" / "st_disjoint" value: +// "bbox" (alias "&&") value: or {"bbox":[minx,miny,maxx,maxy],"srid":4326} +func BuildSpatialCondition(column, operator string, value any) (query string, args []interface{}, ok bool) { + operator = strings.ToLower(strings.TrimSpace(operator)) + + if fn, isPred := spatialPredicates[operator]; isPred { + expr, arg, err := geomArgExpr(value) + if err != nil { + return "", nil, false + } + return fmt.Sprintf("%s(%s, %s)", fn, column, expr), []interface{}{arg}, true + } + + switch operator { + case "st_dwithin": + m, mok := value.(map[string]any) + if !mok { + return "", nil, false + } + expr, arg, err := geomArgExpr(m["geom"]) + if err != nil { + return "", nil, false + } + dist, dok := toFloat(m["distance"]) + if !dok { + return "", nil, false + } + return fmt.Sprintf("ST_DWithin(%s, %s, ?)", column, expr), []interface{}{arg, dist}, true + + case "bbox", "&&": + if m, mok := value.(map[string]any); mok { + if bboxRaw, has := m["bbox"]; has { + coords, cok := toFloatSlice(bboxRaw) + if !cok || len(coords) != 4 { + return "", nil, false + } + srid := 4326 + if s, sok := toFloat(m["srid"]); sok { + srid = int(s) + } + return fmt.Sprintf("%s && ST_MakeEnvelope(?, ?, ?, ?, ?)", column), + []interface{}{coords[0], coords[1], coords[2], coords[3], srid}, true + } + // fall through: treat the map as a GeoJSON geometry + } + expr, arg, err := geomArgExpr(value) + if err != nil { + return "", nil, false + } + return fmt.Sprintf("%s && %s", column, expr), []interface{}{arg}, true + } + + return "", nil, false +} + +// geomArgExpr inspects a geometry value and returns the SQL placeholder +// expression that turns a bound argument into a geometry, plus that argument. +// +// GeoJSON object -> "ST_GeomFromGeoJSON(?)", +// hex EWKB -> "?::geometry", +// WKT / EWKT -> "ST_GeomFromEWKT(?)", +func geomArgExpr(value any) (expr string, arg any, err error) { + switch v := value.(type) { + case nil: + return "", nil, fmt.Errorf("geometry: nil value") + case map[string]any: + b, mErr := json.Marshal(v) + if mErr != nil { + return "", nil, mErr + } + return "ST_GeomFromGeoJSON(?)", string(b), nil + case json.RawMessage: + return "ST_GeomFromGeoJSON(?)", string(v), nil + case []byte: + return geomArgExpr(string(v)) + case string: + s := strings.TrimSpace(v) + if s == "" { + return "", nil, fmt.Errorf("geometry: empty value") + } + if strings.HasPrefix(s, "{") { + return "ST_GeomFromGeoJSON(?)", s, nil + } + if isHexString(s) { + return "?::geometry", s, nil + } + return "ST_GeomFromEWKT(?)", s, nil + default: + return "", nil, fmt.Errorf("geometry: unsupported type %T", value) + } +} + +func isHexString(s string) bool { + if len(s) < 10 || len(s)%2 != 0 { + return false + } + _, err := hex.DecodeString(s) + return err == nil +} + +func toFloat(v any) (float64, bool) { + switch n := v.(type) { + case float64: + return n, true + case float32: + return float64(n), true + case int: + return float64(n), true + case int64: + return float64(n), true + case json.Number: + f, err := n.Float64() + return f, err == nil + case string: + f, err := strconv.ParseFloat(strings.TrimSpace(n), 64) + return f, err == nil + default: + return 0, false + } +} + +func toFloatSlice(v any) ([]float64, bool) { + switch s := v.(type) { + case []float64: + return s, true + case []any: + out := make([]float64, len(s)) + for i, e := range s { + f, ok := toFloat(e) + if !ok { + return nil, false + } + out[i] = f + } + return out, true + default: + return nil, false + } +} + +// IsSpatialOperator reports whether op is a spatial filter operator handled by +// BuildSpatialCondition. +func IsSpatialOperator(op string) bool { + op = strings.ToLower(strings.TrimSpace(op)) + if _, ok := spatialPredicates[op]; ok { + return true + } + return op == "st_dwithin" || op == "bbox" || op == "&&" +} + +// IsVectorOperator reports whether op is a vector similarity filter operator +// handled by BuildVectorCondition. +func IsVectorOperator(op string) bool { + switch strings.ToLower(strings.TrimSpace(op)) { + case "l2_within", "l2distance_within", "euclidean_within", + "cosine_within", "cosinedistance_within", + "ip_within", "inner_within", "negativeinnerproduct_within": + return true + default: + return false + } +} diff --git a/pkg/common/spatial_helpers_test.go b/pkg/common/spatial_helpers_test.go new file mode 100644 index 0000000..830fe39 --- /dev/null +++ b/pkg/common/spatial_helpers_test.go @@ -0,0 +1,145 @@ +package common + +import ( + "testing" +) + +func TestVectorOperator(t *testing.T) { + cases := map[string]string{ + "": "<->", "l2": "<->", "euclidean": "<->", + "cosine": "<=>", "cos": "<=>", + "ip": "<#>", "inner": "<#>", "dot": "<#>", + } + for in, want := range cases { + if got := VectorOperator(in); got != want { + t.Errorf("VectorOperator(%q) = %q, want %q", in, got, want) + } + } +} + +func TestVectorLiteral(t *testing.T) { + cases := []struct { + in any + want string + }{ + {[]float32{1, 2, 3}, "[1,2,3]"}, + {[]float64{1.5, -2}, "[1.5,-2]"}, + {[]int{1, 2}, "[1,2]"}, + {[]any{1.0, 2.0}, "[1,2]"}, + {"[4,5,6]", "[4,5,6]"}, + } + for _, c := range cases { + got, err := VectorLiteral(c.in) + if err != nil || got != c.want { + t.Errorf("VectorLiteral(%v) = %q, %v; want %q", c.in, got, err, c.want) + } + } + if _, err := VectorLiteral("not-a-vector"); err == nil { + t.Error("expected error for malformed string") + } + if _, err := VectorLiteral(42); err == nil { + t.Error("expected error for unsupported type") + } +} + +func TestBuildVectorCondition(t *testing.T) { + q, args, ok := BuildVectorCondition("embedding", "cosine_within", map[string]any{ + "vector": []any{1.0, 2.0, 3.0}, "distance": 0.5, + }) + if !ok { + t.Fatal("expected ok") + } + if q != "embedding <=> ? < ?" { + t.Errorf("query = %q", q) + } + if len(args) != 2 || args[0] != "[1,2,3]" || args[1] != 0.5 { + t.Errorf("args = %v", args) + } + + // explicit comparator + q, _, ok = BuildVectorCondition("v", "l2_within", map[string]any{ + "vector": []float32{1}, "lte": 2.0, + }) + if !ok || q != "v <-> ? <= ?" { + t.Errorf("lte: q=%q ok=%v", q, ok) + } + + // unknown operator + if _, _, ok := BuildVectorCondition("v", "bogus", map[string]any{}); ok { + t.Error("expected not ok for unknown operator") + } + // missing threshold + if _, _, ok := BuildVectorCondition("v", "l2_within", map[string]any{"vector": []float32{1}}); ok { + t.Error("expected not ok without threshold") + } +} + +func TestBuildSpatialCondition_Predicates(t *testing.T) { + q, args, ok := BuildSpatialCondition("geom", "st_intersects", "SRID=4326;POINT(0 0)") + if !ok { + t.Fatal("expected ok") + } + if q != "ST_Intersects(geom, ST_GeomFromEWKT(?))" { + t.Errorf("query = %q", q) + } + if len(args) != 1 || args[0] != "SRID=4326;POINT(0 0)" { + t.Errorf("args = %v", args) + } + + // GeoJSON value + q, args, ok = BuildSpatialCondition("geom", "st_contains", map[string]any{ + "type": "Point", "coordinates": []any{1.0, 2.0}, + }) + if !ok || q != "ST_Contains(geom, ST_GeomFromGeoJSON(?))" { + t.Errorf("geojson: q=%q ok=%v", q, ok) + } + if len(args) != 1 { + t.Errorf("args = %v", args) + } +} + +func TestBuildSpatialCondition_DWithin(t *testing.T) { + q, args, ok := BuildSpatialCondition("geom", "st_dwithin", map[string]any{ + "geom": "SRID=4326;POINT(0 0)", "distance": 1000.0, + }) + if !ok { + t.Fatal("expected ok") + } + if q != "ST_DWithin(geom, ST_GeomFromEWKT(?), ?)" { + t.Errorf("query = %q", q) + } + if len(args) != 2 || args[1] != 1000.0 { + t.Errorf("args = %v", args) + } +} + +func TestBuildSpatialCondition_BBox(t *testing.T) { + q, args, ok := BuildSpatialCondition("geom", "bbox", map[string]any{ + "bbox": []any{0.0, 0.0, 10.0, 10.0}, "srid": 4326.0, + }) + if !ok { + t.Fatal("expected ok") + } + if q != "geom && ST_MakeEnvelope(?, ?, ?, ?, ?)" { + t.Errorf("query = %q", q) + } + if len(args) != 5 || args[4] != 4326 { + t.Errorf("args = %v", args) + } +} + +func TestIsSpatialAndVectorOperator(t *testing.T) { + for _, op := range []string{"st_dwithin", "st_intersects", "bbox", "&&"} { + if !IsSpatialOperator(op) { + t.Errorf("%q should be spatial", op) + } + } + for _, op := range []string{"l2_within", "cosine_within", "ip_within"} { + if !IsVectorOperator(op) { + t.Errorf("%q should be vector", op) + } + } + if IsSpatialOperator("eq") || IsVectorOperator("eq") { + t.Error("eq is neither spatial nor vector") + } +} diff --git a/pkg/common/types.go b/pkg/common/types.go index 4a4a683..92295d6 100644 --- a/pkg/common/types.go +++ b/pkg/common/types.go @@ -42,6 +42,10 @@ type RequestOptions struct { CursorBackward string `json:"cursor_backward"` FetchRowNumber *string `json:"fetch_row_number"` + // VectorSearch performs a pgvector nearest-neighbour ordering (KNN) and + // optionally returns the computed distance as an extra column. + VectorSearch *VectorSearchOption `json:"vector_search"` + // Join table aliases (used for validation of prefixed columns in filters/sorts) // Not serialized to JSON as it's internal validation state JoinAliases []string `json:"-"` @@ -114,6 +118,17 @@ func ResolveSortColumns(sort []SortOption, pkName string) []SortOption { return resolved } +// VectorSearchOption describes a pgvector KNN search: order rows by the distance +// between Column and Vector using Metric, and (when As is set) select that +// distance as an additional result column. +type VectorSearchOption struct { + Column string `json:"column"` + Vector []float32 `json:"vector"` + Metric string `json:"metric"` // "l2" (default) | "cosine" | "ip" + As string `json:"as"` // distance column alias; default "_distance" + Direction string `json:"direction"` // "asc" (default) | "desc" +} + type CustomOperator struct { Name string `json:"name"` SQL string `json:"sql"` diff --git a/pkg/openapi/generator.go b/pkg/openapi/generator.go index 9acc102..ee73b5f 100644 --- a/pkg/openapi/generator.go +++ b/pkg/openapi/generator.go @@ -7,6 +7,7 @@ import ( "strings" "github.com/bitechdev/ResolveSpec/pkg/modelregistry" + "github.com/bitechdev/ResolveSpec/pkg/spectypes" ) // OpenAPISpec represents the OpenAPI 3.0 specification structure @@ -440,6 +441,28 @@ func (g *Generator) generatePropertySchema(field reflect.StructField) *Schema { schema.Description = desc } + // spectypes PostGIS / pgvector wrappers get dedicated schemas. + if n, ok := spectypes.SQLTypeName(field.Type); ok { + switch n { + case "geometry", "geography": + schema.Type = "object" + schema.Format = "geojson" + return schema + case "vector", "halfvec": + schema.Type = "array" + schema.Items = &Schema{Type: "number"} + schema.Format = "vector" + return schema + case "sparsevec": + schema.Type = "object" + schema.Format = "sparsevec" + return schema + case "bit": + schema.Type = "string" + return schema + } + } + switch fieldType.Kind() { case reflect.String: schema.Type = "string" diff --git a/pkg/reflection/model_utils_geo_test.go b/pkg/reflection/model_utils_geo_test.go new file mode 100644 index 0000000..802a70b --- /dev/null +++ b/pkg/reflection/model_utils_geo_test.go @@ -0,0 +1,56 @@ +package reflection + +import ( + "testing" + + "github.com/bitechdev/ResolveSpec/pkg/spectypes" +) + +type geoModel struct { + ID int64 `json:"id"` + Location spectypes.SqlGeometry `json:"location"` + Area spectypes.SqlGeography `json:"area"` + Embedding spectypes.SqlVector `json:"embedding"` + HalfEmb spectypes.SqlHalfVector `json:"half_emb"` + Name spectypes.SqlString `json:"name"` +} + +func TestGetColumnSQLTypeName(t *testing.T) { + m := geoModel{} + cases := map[string]string{ + "location": "geometry", + "area": "geography", + "embedding": "vector", + "half_emb": "halfvec", + "name": "text", + } + for col, want := range cases { + got, ok := GetColumnSQLTypeName(m, col) + if !ok || got != want { + t.Errorf("GetColumnSQLTypeName(%q) = %q, %v; want %q", col, got, ok, want) + } + } + if _, ok := GetColumnSQLTypeName(m, "id"); ok { + t.Error("id is not a spectypes column") + } +} + +func TestIsSpatialColumn(t *testing.T) { + m := geoModel{} + if !IsSpatialColumn(m, "location") || !IsSpatialColumn(m, "area") { + t.Error("location/area should be spatial") + } + if IsSpatialColumn(m, "embedding") || IsSpatialColumn(m, "name") { + t.Error("embedding/name should not be spatial") + } +} + +func TestIsVectorColumn(t *testing.T) { + m := geoModel{} + if !IsVectorColumn(m, "embedding") || !IsVectorColumn(m, "half_emb") { + t.Error("embedding/half_emb should be vector") + } + if IsVectorColumn(m, "location") || IsVectorColumn(m, "name") { + t.Error("location/name should not be vector") + } +} diff --git a/pkg/reflection/spectypes_helpers.go b/pkg/reflection/spectypes_helpers.go new file mode 100644 index 0000000..a7f3a3e --- /dev/null +++ b/pkg/reflection/spectypes_helpers.go @@ -0,0 +1,94 @@ +package reflection + +import ( + "reflect" + + "github.com/bitechdev/ResolveSpec/pkg/spectypes" +) + +// getColumnFieldType resolves the reflect.Type of the struct field that backs +// colName (matched by json tag, field name or snake_case), following the same +// rules as GetColumnTypeFromModel. +func getColumnFieldType(model interface{}, colName string) (reflect.Type, bool) { + if model == nil { + return nil, false + } + sourceColName := ExtractSourceColumn(colName) + + modelType := reflect.TypeOf(model) + for modelType != nil && modelType.Kind() == reflect.Pointer { + modelType = modelType.Elem() + } + if modelType == nil || modelType.Kind() != reflect.Struct { + return nil, false + } + + for i := 0; i < modelType.NumField(); i++ { + field := modelType.Field(i) + + if jsonTag := field.Tag.Get("json"); jsonTag != "" { + if name := jsonTagName(jsonTag); name == sourceColName { + return field.Type, true + } + } + if equalFold(field.Name, sourceColName) { + return field.Type, true + } + if ToSnakeCase(field.Name) == sourceColName { + return field.Type, true + } + } + return nil, false +} + +func jsonTagName(tag string) string { + for i := 0; i < len(tag); i++ { + if tag[i] == ',' { + return tag[:i] + } + } + return tag +} + +func equalFold(a, b string) bool { + if len(a) != len(b) { + return false + } + for i := 0; i < len(a); i++ { + ca, cb := a[i], b[i] + if 'A' <= ca && ca <= 'Z' { + ca += 'a' - 'A' + } + if 'A' <= cb && cb <= 'Z' { + cb += 'a' - 'A' + } + if ca != cb { + return false + } + } + return true +} + +// GetColumnSQLTypeName returns the canonical PostgreSQL type name for a column +// backed by a spectypes wrapper (e.g. "geometry", "vector", "jsonb"), or +// ("", false) if the column is not found or not a spectypes type. +func GetColumnSQLTypeName(model interface{}, colName string) (string, bool) { + t, ok := getColumnFieldType(model, colName) + if !ok { + return "", false + } + return spectypes.SQLTypeName(t) +} + +// IsSpatialColumn reports whether colName is backed by a PostGIS +// geometry/geography wrapper. +func IsSpatialColumn(model interface{}, colName string) bool { + t, ok := getColumnFieldType(model, colName) + return ok && spectypes.IsSpatialType(t) +} + +// IsVectorColumn reports whether colName is backed by a pgvector wrapper. +func IsVectorColumn(model interface{}, colName string) bool { + t, ok := getColumnFieldType(model, colName) + return ok && spectypes.IsVectorType(t) +} diff --git a/pkg/resolvespec/filter_test.go b/pkg/resolvespec/filter_test.go index 938efde..6ee2ea7 100644 --- a/pkg/resolvespec/filter_test.go +++ b/pkg/resolvespec/filter_test.go @@ -87,6 +87,43 @@ func TestBuildFilterCondition(t *testing.T) { expectedCondition: "", expectedArgsCount: 0, }, + { + name: "st_dwithin spatial operator", + filter: common.FilterOption{ + Column: "geom", + Operator: "st_dwithin", + Value: map[string]any{ + "geom": "SRID=4326;POINT(0 0)", + "distance": 1000.0, + }, + }, + expectedCondition: "ST_DWithin(geom, ST_GeomFromEWKT(?), ?)", + expectedArgsCount: 2, + }, + { + name: "st_intersects spatial operator", + filter: common.FilterOption{ + Column: "geom", + Operator: "st_intersects", + Value: "SRID=4326;POLYGON((0 0,1 0,1 1,0 1,0 0))", + LogicOperator: "", + }, + expectedCondition: "ST_Intersects(geom, ST_GeomFromEWKT(?))", + expectedArgsCount: 1, + }, + { + name: "l2_within vector operator", + filter: common.FilterOption{ + Column: "embedding", + Operator: "l2_within", + Value: map[string]any{ + "vector": []any{1.0, 2.0, 3.0}, + "distance": 0.5, + }, + }, + expectedCondition: "embedding <-> ? < ?", + expectedArgsCount: 2, + }, } for _, tt := range tests { diff --git a/pkg/resolvespec/handler.go b/pkg/resolvespec/handler.go index 60d8629..5772683 100644 --- a/pkg/resolvespec/handler.go +++ b/pkg/resolvespec/handler.go @@ -16,6 +16,7 @@ import ( "github.com/bitechdev/ResolveSpec/pkg/common" "github.com/bitechdev/ResolveSpec/pkg/logger" "github.com/bitechdev/ResolveSpec/pkg/reflection" + "github.com/bitechdev/ResolveSpec/pkg/spectypes" ) // FallbackHandler is a function that handles requests when no model is found @@ -332,8 +333,13 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st query = query.Table(tableName) } - if len(options.Columns) == 0 && (len(options.ComputedColumns) > 0) { - logger.Debug("Populating options.Columns with all model columns since computed columns are additions") + vectorSearchActive := options.VectorSearch != nil && + options.VectorSearch.Column != "" && len(options.VectorSearch.Vector) > 0 + + if len(options.Columns) == 0 && + (len(options.ComputedColumns) > 0 || + (vectorSearchActive && options.VectorSearch.As != "")) { + logger.Debug("Populating options.Columns with all model columns since computed/vector columns are additions") options.Columns = reflection.GetSQLModelColumns(model) } @@ -352,6 +358,29 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st } } + // pgvector KNN search: order by distance to the query vector and, + // optionally, return that distance as an extra column. Postgres only. + if vectorSearchActive { + vs := options.VectorSearch + op := common.VectorOperator(vs.Metric) + lit, litErr := common.VectorLiteral(vs.Vector) + if litErr != nil { + logger.Error("Invalid vector search vector: %v", litErr) + statusCode, errCode, errMsg = http.StatusBadRequest, "invalid_vector_search", "Invalid vector search vector" + return litErr + } + col := common.QuoteIdent(vs.Column) + dir := "ASC" + if strings.EqualFold(vs.Direction, "desc") { + dir = "DESC" + } + if vs.As != "" { + query = query.ColumnExpr(fmt.Sprintf("(%s %s ?) AS %s", col, op, common.QuoteIdent(vs.As)), lit) + } + query = query.OrderExpr(fmt.Sprintf("%s %s ? %s", col, op, dir), lit) + logger.Debug("Applying vector search on %s (%s)", vs.Column, op) + } + // Apply preloading if len(options.Preload) > 0 { var err error @@ -1909,7 +1938,21 @@ func (h *Handler) buildFilterCondition(filter common.FilterOption) (conditionStr return "", nil } default: - return "", nil + if common.IsSpatialOperator(filter.Operator) { + q, a, ok := common.BuildSpatialCondition(filter.Column, filter.Operator, filter.Value) + if !ok { + return "", nil + } + condition, args = q, a + } else if common.IsVectorOperator(filter.Operator) { + q, a, ok := common.BuildVectorCondition(filter.Column, filter.Operator, filter.Value) + if !ok { + return "", nil + } + condition, args = q, a + } else { + return "", nil + } } return condition, args @@ -1958,7 +2001,21 @@ func (h *Handler) applyFilter(query common.SelectQuery, filter common.FilterOpti return query } default: - return query + if common.IsSpatialOperator(filter.Operator) { + q, a, ok := common.BuildSpatialCondition(filter.Column, filter.Operator, filter.Value) + if !ok { + return query + } + condition, args = q, a + } else if common.IsVectorOperator(filter.Operator) { + q, a, ok := common.BuildVectorCondition(filter.Column, filter.Operator, filter.Value) + if !ok { + return query + } + condition, args = q, a + } else { + return query + } } // Apply filter with appropriate logic operator @@ -2094,9 +2151,20 @@ func (h *Handler) generateMetadata(schema, entity string, model interface{}) *co continue } + colTypeStr := getColumnType(columnField) + // Fill the gap for spectypes wrappers whose Go representation (struct or + // slice) has no obvious SQL mapping — PostGIS geometry/geography and + // pgvector vector/halfvec/sparsevec/bit. A gorm `type:` tag still wins + // (dimensioned types like vector(1536)). + if colTypeStr == "unknown" && !strings.Contains(field.Tag.Get("gorm"), "type:") { + if n, ok := spectypes.SQLTypeName(field.Type); ok { + colTypeStr = n + } + } + column := common.Column{ Name: jsonName, - Type: getColumnType(columnField), + Type: colTypeStr, IsNullable: isSQLType || isNullable(field), IsPrimary: strings.Contains(gormTag, "primaryKey"), IsUnique: strings.Contains(gormTag, "unique") || strings.Contains(gormTag, "uniqueIndex"), diff --git a/pkg/restheadspec/handler.go b/pkg/restheadspec/handler.go index c2c9018..e43a131 100644 --- a/pkg/restheadspec/handler.go +++ b/pkg/restheadspec/handler.go @@ -419,7 +419,12 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st // If we have computed columns/expressions but options.Columns is empty, // populate it with all model columns first since computed columns are additions - if len(options.Columns) == 0 && (len(options.ComputedQL) > 0 || len(options.ComputedColumns) > 0) { + vectorSearchActive := options.VectorSearch != nil && + options.VectorSearch.Column != "" && len(options.VectorSearch.Vector) > 0 + + if len(options.Columns) == 0 && + (len(options.ComputedQL) > 0 || len(options.ComputedColumns) > 0 || + (vectorSearchActive && options.VectorSearch.As != "")) { logger.Debug("Populating options.Columns with all model columns since computed columns are additions") options.Columns = reflection.GetSQLModelColumns(model) } @@ -472,6 +477,29 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st } + // pgvector KNN search: order by distance to the query vector and, + // optionally, return that distance as an extra column. Postgres only. + if vectorSearchActive { + vs := options.VectorSearch + op := common.VectorOperator(vs.Metric) + lit, litErr := common.VectorLiteral(vs.Vector) + if litErr != nil { + logger.Error("Invalid vector search vector: %v", litErr) + statusCode, errCode, errMsg = http.StatusBadRequest, "invalid_vector_search", "Invalid vector search vector" + return litErr + } + col := common.QuoteIdent(vs.Column) + dir := "ASC" + if strings.EqualFold(vs.Direction, "desc") { + dir = "DESC" + } + if vs.As != "" { + query = query.ColumnExpr(fmt.Sprintf("(%s %s ?) AS %s", col, op, common.QuoteIdent(vs.As)), lit) + } + query = query.OrderExpr(fmt.Sprintf("%s %s ? %s", col, op, dir), lit) + logger.Debug("Applying vector search on %s (%s)", vs.Column, op) + } + // Apply expand (Just expand to Preload for now) for _, expand := range options.Expand { logger.Debug("Applying expand: %s", expand.Relation) @@ -2330,6 +2358,18 @@ func (h *Handler) applyFilter(query common.SelectQuery, filter common.FilterOpti colName := h.qualifyColumnName(filter.Column, tableName) return applyWhere(fmt.Sprintf("(%s IS NOT NULL AND %s != '')", colName, colName)) default: + if common.IsSpatialOperator(filter.Operator) { + if cond, sargs, ok := common.BuildSpatialCondition(rawQualifiedColumn, filter.Operator, filter.Value); ok { + return applyWhere(cond, sargs...) + } + return query + } + if common.IsVectorOperator(filter.Operator) { + if cond, vargs, ok := common.BuildVectorCondition(rawQualifiedColumn, filter.Operator, filter.Value); ok { + return applyWhere(cond, vargs...) + } + return query + } logger.Warn("Unknown filter operator: %s, defaulting to equals", filter.Operator) return applyWhere(fmt.Sprintf("%s = ?", qualifiedColumn), filter.Value) } @@ -2429,6 +2469,20 @@ func (h *Handler) buildFilterCondition(qualifiedColumn string, filter *common.Fi colName := h.qualifyColumnName(filter.Column, tableName) return fmt.Sprintf("(%s IS NOT NULL AND %s != '')", colName, colName), nil default: + if common.IsSpatialOperator(filter.Operator) { + rawCol := h.qualifyColumnName(filter.Column, tableName) + if cond, sargs, ok := common.BuildSpatialCondition(rawCol, filter.Operator, filter.Value); ok { + return cond, sargs + } + return "", nil + } + if common.IsVectorOperator(filter.Operator) { + rawCol := h.qualifyColumnName(filter.Column, tableName) + if cond, vargs, ok := common.BuildVectorCondition(rawCol, filter.Operator, filter.Value); ok { + return cond, vargs + } + return "", nil + } logger.Warn("Unknown filter operator: %s, defaulting to equals", filter.Operator) return fmt.Sprintf("%s = ?", qualifiedColumn), []interface{}{filter.Value} } diff --git a/pkg/restheadspec/headers.go b/pkg/restheadspec/headers.go index 536e1af..35e70eb 100644 --- a/pkg/restheadspec/headers.go +++ b/pkg/restheadspec/headers.go @@ -182,6 +182,22 @@ func (h *Handler) parseOptionsFromHeaders(r common.Request, model interface{}) E h.parseSearchOp(&options, key, decodedValue, "AND") case strings.HasPrefix(key, "x-searchcols"): options.SearchColumns = h.parseCommaSeparated(decodedValue) + case strings.HasPrefix(key, "x-spatialfilter-"): + h.parseGeoFilter(&options, key, "x-spatialfilter-", decodedValue) + case strings.HasPrefix(key, "x-vectorfilter-"): + h.parseGeoFilter(&options, key, "x-vectorfilter-", decodedValue) + + // pgvector KNN search + case key == "x-vector-search-vector": + h.ensureVectorSearch(&options).Vector = parseFloat32List(decodedValue) + case key == "x-vector-search-as": + h.ensureVectorSearch(&options).As = decodedValue + case key == "x-vector-search-dir": + h.ensureVectorSearch(&options).Direction = decodedValue + case strings.HasPrefix(key, "x-vector-search-"): + vs := h.ensureVectorSearch(&options) + vs.Column = strings.TrimPrefix(key, "x-vector-search-") + vs.Metric = decodedValue case strings.HasPrefix(key, "x-custom-sql-w"): if options.CustomSQLWhere != "" { options.CustomSQLWhere = fmt.Sprintf("%s AND (%s)", options.CustomSQLWhere, decodedValue) @@ -309,6 +325,83 @@ func (h *Handler) parseOptionsFromHeaders(r common.Request, model interface{}) E return options } +// ensureVectorSearch returns the options' VectorSearchOption, allocating it on +// first use. +func (h *Handler) ensureVectorSearch(options *ExtendedRequestOptions) *common.VectorSearchOption { + if options.VectorSearch == nil { + options.VectorSearch = &common.VectorSearchOption{} + } + return options.VectorSearch +} + +// parseFloat32List parses a JSON array ("[1,2,3]") or comma-separated list into +// a []float32. +func parseFloat32List(value string) []float32 { + value = strings.TrimSpace(value) + if value == "" { + return nil + } + var raw []float64 + if err := json.Unmarshal([]byte(value), &raw); err == nil { + out := make([]float32, len(raw)) + for i, f := range raw { + out[i] = float32(f) + } + return out + } + parts := strings.Split(strings.Trim(value, "[]"), ",") + out := make([]float32, 0, len(parts)) + for _, p := range parts { + f, err := strconv.ParseFloat(strings.TrimSpace(p), 32) + if err != nil { + return nil + } + out = append(out, float32(f)) + } + return out +} + +// parseGeoFilter parses an x-spatialfilter- / x-vectorfilter- header. +// The value is a JSON object: {"op":"st_dwithin","geom":...,"distance":...} or +// {"op":"st_intersects","value":}. An optional "logic":"or" controls +// how the filter combines with the previous one. +func (h *Handler) parseGeoFilter(options *ExtendedRequestOptions, key, prefix, value string) { + col := strings.TrimPrefix(key, prefix) + if col == "" || strings.TrimSpace(value) == "" { + return + } + var raw map[string]interface{} + if err := json.Unmarshal([]byte(value), &raw); err != nil { + logger.Warn("Invalid %s%s filter JSON: %v", prefix, col, err) + return + } + op, _ := raw["op"].(string) + if op == "" { + logger.Warn("%s%s filter missing \"op\"", prefix, col) + return + } + logicOp := "AND" + if lo, ok := raw["logic"].(string); ok && strings.EqualFold(lo, "or") { + logicOp = "OR" + } + + var fv interface{} + if v, ok := raw["value"]; ok { + fv = v + } else { + delete(raw, "op") + delete(raw, "logic") + fv = raw + } + + options.Filters = append(options.Filters, common.FilterOption{ + Column: col, + Operator: op, + Value: fv, + LogicOperator: logicOp, + }) +} + // parseSelectFields parses x-select-fields header func (h *Handler) parseSelectFields(options *ExtendedRequestOptions, value string) { if value == "" { @@ -1365,6 +1458,14 @@ func (h *Handler) ValidateAndAdjustFilterForColumnType(filter *common.FilterOpti return ColumnCastInfo{NeedsCast: false, IsNumericType: false} } + // Never cast geometry/geography or pgvector columns to TEXT — spatial and + // vector operators need the native column type. Also bypass when the + // operator itself is spatial/vector (e.g. st_dwithin, l2_within). + if common.IsSpatialOperator(filter.Operator) || common.IsVectorOperator(filter.Operator) || + reflection.IsSpatialColumn(model, filter.Column) || reflection.IsVectorColumn(model, filter.Column) { + return ColumnCastInfo{NeedsCast: false, IsNumericType: false} + } + colType := reflection.GetColumnTypeFromModel(model, filter.Column) if colType == reflect.Invalid { // Column not found in model, no casting needed diff --git a/pkg/restheadspec/spatial_filter_test.go b/pkg/restheadspec/spatial_filter_test.go new file mode 100644 index 0000000..8478f9c --- /dev/null +++ b/pkg/restheadspec/spatial_filter_test.go @@ -0,0 +1,124 @@ +package restheadspec + +import ( + "testing" + + "github.com/bitechdev/ResolveSpec/pkg/common" +) + +func TestBuildFilterCondition_Spatial(t *testing.T) { + h := &Handler{} + + tests := []struct { + name string + filter common.FilterOption + wantCond string + wantCount int + }{ + { + name: "st_dwithin", + filter: common.FilterOption{ + Column: "geom", + Operator: "st_dwithin", + Value: map[string]interface{}{ + "geom": "SRID=4326;POINT(0 0)", + "distance": 1000.0, + }, + }, + wantCond: "ST_DWithin(geom, ST_GeomFromEWKT(?), ?)", + wantCount: 2, + }, + { + name: "st_intersects", + filter: common.FilterOption{ + Column: "geom", + Operator: "st_intersects", + Value: "SRID=4326;POINT(0 0)", + }, + wantCond: "ST_Intersects(geom, ST_GeomFromEWKT(?))", + wantCount: 1, + }, + { + name: "l2_within", + filter: common.FilterOption{ + Column: "embedding", + Operator: "l2_within", + Value: map[string]interface{}{ + "vector": []interface{}{1.0, 2.0}, + "distance": 0.3, + }, + }, + wantCond: "embedding <-> ? < ?", + wantCount: 2, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + f := tt.filter + cond, args := h.buildFilterCondition(f.Column, &f, "") + if cond != tt.wantCond { + t.Errorf("cond = %q, want %q", cond, tt.wantCond) + } + if len(args) != tt.wantCount { + t.Errorf("args = %d, want %d", len(args), tt.wantCount) + } + }) + } +} + +func TestParseGeoFilter(t *testing.T) { + h := &Handler{} + options := &ExtendedRequestOptions{} + + h.parseGeoFilter(options, "x-spatialfilter-geom", "x-spatialfilter-", + `{"op":"st_dwithin","geom":"SRID=4326;POINT(0 0)","distance":500}`) + + if len(options.Filters) != 1 { + t.Fatalf("expected 1 filter, got %d", len(options.Filters)) + } + f := options.Filters[0] + if f.Column != "geom" || f.Operator != "st_dwithin" { + t.Errorf("filter = %+v", f) + } + m, ok := f.Value.(map[string]interface{}) + if !ok || m["distance"] != float64(500) { + t.Errorf("value = %v", f.Value) + } + if _, has := m["op"]; has { + t.Error("op should be stripped from value map") + } +} + +func TestParseGeoFilter_ExplicitValue(t *testing.T) { + h := &Handler{} + options := &ExtendedRequestOptions{} + + h.parseGeoFilter(options, "x-vectorfilter-embedding", "x-vectorfilter-", + `{"op":"cosine_within","logic":"or","value":{"vector":[1,2,3],"distance":0.2}}`) + + if len(options.Filters) != 1 { + t.Fatalf("expected 1 filter, got %d", len(options.Filters)) + } + f := options.Filters[0] + if f.Operator != "cosine_within" || f.LogicOperator != "OR" { + t.Errorf("filter = %+v", f) + } + if _, ok := f.Value.(map[string]interface{}); !ok { + t.Errorf("value type = %T", f.Value) + } +} + +func TestParseFloat32List(t *testing.T) { + got := parseFloat32List("[1,2.5,3]") + if len(got) != 3 || got[1] != 2.5 { + t.Errorf("json array = %v", got) + } + got = parseFloat32List("1, 2, 3") + if len(got) != 3 || got[2] != 3 { + t.Errorf("csv = %v", got) + } + if parseFloat32List("") != nil { + t.Error("empty should be nil") + } +} diff --git a/pkg/spectypes/sql_geo_types.go b/pkg/spectypes/sql_geo_types.go new file mode 100644 index 0000000..9a9ebae --- /dev/null +++ b/pkg/spectypes/sql_geo_types.go @@ -0,0 +1,239 @@ +package spectypes + +import ( + "database/sql/driver" + "encoding/hex" + "encoding/json" + "fmt" + "strconv" + "strings" +) + +// ── PostGIS geometry / geography ───────────────────────────────────────────── + +// SqlGeometry is a nullable PostGIS `geometry` column. +// +// On Scan it accepts the PostGIS default hex-EWKB text output, a raw GeoJSON +// object, or a WKT/EWKT string (e.g. when the column is selected via +// ST_AsGeoJSON / ST_AsText). Internally it holds a canonical GeoJSON geometry +// object plus the SRID. +// +// On Value it emits `SRID=;` text. PostGIS registers an implicit +// text -> geometry cast, so parameterised inserts/updates work without wrapping +// the placeholder in a constructor function. +// +// MarshalJSON emits the GeoJSON geometry object (or null). +type SqlGeometry struct { + GeoJSON json.RawMessage + SRID int + Valid bool +} + +// SqlGeography is identical to SqlGeometry but maps to a PostGIS `geography` +// column. Coordinates are always lon/lat and the default SRID is 4326. +type SqlGeography struct { + SqlGeometry +} + +func (g *SqlGeometry) Scan(value any) error { + if value == nil { + g.Valid = false + g.GeoJSON = nil + g.SRID = 0 + return nil + } + + var s string + switch v := value.(type) { + case string: + s = v + case []byte: + s = string(v) + default: + return fmt.Errorf("SqlGeometry: cannot scan type %T", value) + } + + s = strings.TrimSpace(s) + if s == "" { + g.Valid = false + g.GeoJSON = nil + return nil + } + + switch { + case strings.HasPrefix(s, "{"): + // GeoJSON object. + if _, err := geoJSONToGeom([]byte(s)); err != nil { + return fmt.Errorf("SqlGeometry: invalid GeoJSON: %w", err) + } + g.GeoJSON = json.RawMessage(s) + g.Valid = true + return nil + case isHex(s): + gj, srid, err := DecodeEWKBHex(s) + if err != nil { + return fmt.Errorf("SqlGeometry: %w", err) + } + g.GeoJSON = gj + g.SRID = srid + g.Valid = true + return nil + default: + // WKT / EWKT text. + srid, wkt := splitEWKT(s) + gj, err := wktToGeoJSON(wkt) + if err != nil { + return fmt.Errorf("SqlGeometry: %w", err) + } + g.GeoJSON = gj + g.SRID = srid + g.Valid = true + return nil + } +} + +func (g SqlGeometry) Value() (driver.Value, error) { + if !g.Valid || len(g.GeoJSON) == 0 { + return nil, nil + } + wkt, err := GeoJSONToWKT(g.GeoJSON) + if err != nil { + return nil, err + } + srid := g.SRID + if srid == 0 { + srid = 4326 + } + return fmt.Sprintf("SRID=%d;%s", srid, wkt), nil +} + +func (g SqlGeometry) MarshalJSON() ([]byte, error) { + if !g.Valid || len(g.GeoJSON) == 0 { + return []byte("null"), nil + } + return g.GeoJSON, nil +} + +func (g *SqlGeometry) UnmarshalJSON(b []byte) error { + s := strings.TrimSpace(string(b)) + if s == "" || s == "null" { + g.Valid = false + g.GeoJSON = nil + g.SRID = 0 + return nil + } + + if strings.HasPrefix(s, "{") { + if _, err := geoJSONToGeom(b); err != nil { + return fmt.Errorf("SqlGeometry: invalid GeoJSON: %w", err) + } + g.GeoJSON = append(json.RawMessage(nil), b...) + g.Valid = true + return nil + } + + // String value: EWKT / WKT / hex-EWKB. + var str string + if err := json.Unmarshal(b, &str); err != nil { + return fmt.Errorf("SqlGeometry: cannot unmarshal %s", b) + } + str = strings.TrimSpace(str) + if str == "" { + g.Valid = false + g.GeoJSON = nil + return nil + } + if isHex(str) { + gj, srid, err := DecodeEWKBHex(str) + if err != nil { + return fmt.Errorf("SqlGeometry: %w", err) + } + g.GeoJSON = gj + g.SRID = srid + g.Valid = true + return nil + } + srid, wkt := splitEWKT(str) + gj, err := wktToGeoJSON(wkt) + if err != nil { + return fmt.Errorf("SqlGeometry: %w", err) + } + g.GeoJSON = gj + g.SRID = srid + g.Valid = true + return nil +} + +// WKT returns the geometry as a plain WKT string (no SRID prefix). +func (g SqlGeometry) WKT() string { + if !g.Valid || len(g.GeoJSON) == 0 { + return "" + } + wkt, err := GeoJSONToWKT(g.GeoJSON) + if err != nil { + return "" + } + return wkt +} + +// EWKT returns the geometry as `SRID=;`. +func (g SqlGeometry) EWKT() string { + wkt := g.WKT() + if wkt == "" { + return "" + } + srid := g.SRID + if srid == 0 { + srid = 4326 + } + return fmt.Sprintf("SRID=%d;%s", srid, wkt) +} + +// NewSqlGeometryFromGeoJSON builds a SqlGeometry from a GeoJSON geometry object. +func NewSqlGeometryFromGeoJSON(geojson []byte, srid int) (SqlGeometry, error) { + if _, err := geoJSONToGeom(geojson); err != nil { + return SqlGeometry{}, err + } + return SqlGeometry{GeoJSON: append(json.RawMessage(nil), geojson...), SRID: srid, Valid: true}, nil +} + +// NewSqlGeometryFromEWKT builds a SqlGeometry from an EWKT or WKT string. +func NewSqlGeometryFromEWKT(ewkt string) (SqlGeometry, error) { + srid, wkt := splitEWKT(strings.TrimSpace(ewkt)) + gj, err := wktToGeoJSON(wkt) + if err != nil { + return SqlGeometry{}, err + } + return SqlGeometry{GeoJSON: gj, SRID: srid, Valid: true}, nil +} + +// ── helpers ───────────────────────────────────────────────────────────────── + +func isHex(s string) bool { + if len(s) < 10 || len(s)%2 != 0 { + return false + } + _, err := hex.DecodeString(s) + return err == nil +} + +// splitEWKT separates an optional `SRID=;` prefix from a WKT body. +func splitEWKT(s string) (srid int, wkt string) { + if strings.HasPrefix(strings.ToUpper(s), "SRID=") { + if idx := strings.Index(s, ";"); idx > 0 { + if n, err := strconv.Atoi(strings.TrimSpace(s[5:idx])); err == nil { + return n, strings.TrimSpace(s[idx+1:]) + } + } + } + return 0, s +} + +// wktToGeoJSON parses a (subset of) WKT into a GeoJSON geometry object. +func wktToGeoJSON(wkt string) ([]byte, error) { + g, err := parseWKT(wkt) + if err != nil { + return nil, err + } + return geomToGeoJSON(g) +} diff --git a/pkg/spectypes/sql_geo_types_test.go b/pkg/spectypes/sql_geo_types_test.go new file mode 100644 index 0000000..a58a9ad --- /dev/null +++ b/pkg/spectypes/sql_geo_types_test.go @@ -0,0 +1,126 @@ +package spectypes + +import ( + "encoding/json" + "testing" +) + +// SRID=4326;POINT (1 2) +const pointHexEWKB = "0101000020E6100000000000000000F03F0000000000000040" + +func TestSqlGeometry_ScanHexEWKB(t *testing.T) { + var g SqlGeometry + if err := g.Scan(pointHexEWKB); err != nil { + t.Fatalf("Scan: %v", err) + } + if !g.Valid || g.SRID != 4326 { + t.Fatalf("got Valid=%v SRID=%d", g.Valid, g.SRID) + } + if !jsonEqual(t, g.GeoJSON, `{"type":"Point","coordinates":[1,2]}`) { + t.Errorf("GeoJSON = %s", g.GeoJSON) + } +} + +func TestSqlGeometry_ScanGeoJSON(t *testing.T) { + var g SqlGeometry + if err := g.Scan(`{"type":"Point","coordinates":[3,4]}`); err != nil { + t.Fatalf("Scan: %v", err) + } + if !g.Valid { + t.Fatal("expected valid") + } + if !jsonEqual(t, g.GeoJSON, `{"type":"Point","coordinates":[3,4]}`) { + t.Errorf("GeoJSON = %s", g.GeoJSON) + } +} + +func TestSqlGeometry_ScanEWKT(t *testing.T) { + var g SqlGeometry + if err := g.Scan("SRID=3857;POINT (5 6)"); err != nil { + t.Fatalf("Scan: %v", err) + } + if g.SRID != 3857 { + t.Errorf("SRID = %d, want 3857", g.SRID) + } + if !jsonEqual(t, g.GeoJSON, `{"type":"Point","coordinates":[5,6]}`) { + t.Errorf("GeoJSON = %s", g.GeoJSON) + } +} + +func TestSqlGeometry_Value(t *testing.T) { + g, err := NewSqlGeometryFromGeoJSON([]byte(`{"type":"Point","coordinates":[1,2]}`), 4326) + if err != nil { + t.Fatalf("New: %v", err) + } + v, err := g.Value() + if err != nil { + t.Fatalf("Value: %v", err) + } + if v != "SRID=4326;POINT (1 2)" { + t.Errorf("Value = %v, want SRID=4326;POINT (1 2)", v) + } +} + +func TestSqlGeometry_ValueDefaultsSRID(t *testing.T) { + g, _ := NewSqlGeometryFromGeoJSON([]byte(`{"type":"Point","coordinates":[1,2]}`), 0) + v, _ := g.Value() + if v != "SRID=4326;POINT (1 2)" { + t.Errorf("Value = %v", v) + } +} + +func TestSqlGeometry_JSON(t *testing.T) { + g, _ := NewSqlGeometryFromEWKT("SRID=4326;POINT (1 2)") + b, err := json.Marshal(g) + if err != nil { + t.Fatalf("Marshal: %v", err) + } + if !jsonEqual(t, b, `{"type":"Point","coordinates":[1,2]}`) { + t.Errorf("json = %s", b) + } + + var back SqlGeometry + if err := json.Unmarshal([]byte(`{"type":"Point","coordinates":[7,8]}`), &back); err != nil { + t.Fatalf("Unmarshal: %v", err) + } + if !back.Valid || !jsonEqual(t, back.GeoJSON, `{"type":"Point","coordinates":[7,8]}`) { + t.Errorf("unmarshal = %+v", back) + } + + // Unmarshal also accepts an EWKT string. + var fromStr SqlGeometry + if err := json.Unmarshal([]byte(`"SRID=4326;POINT(9 10)"`), &fromStr); err != nil { + t.Fatalf("Unmarshal string: %v", err) + } + if !jsonEqual(t, fromStr.GeoJSON, `{"type":"Point","coordinates":[9,10]}`) { + t.Errorf("fromStr = %s", fromStr.GeoJSON) + } +} + +func TestSqlGeometry_Null(t *testing.T) { + var g SqlGeometry + if err := g.Scan(nil); err != nil { + t.Fatalf("Scan(nil): %v", err) + } + if g.Valid { + t.Error("expected invalid") + } + v, err := g.Value() + if err != nil || v != nil { + t.Errorf("Value = %v, %v", v, err) + } + b, _ := json.Marshal(g) + if string(b) != "null" { + t.Errorf("json = %s", b) + } +} + +func TestSqlGeography_Embeds(t *testing.T) { + var g SqlGeography + if err := g.Scan("SRID=4326;POINT (1 2)"); err != nil { + t.Fatalf("Scan: %v", err) + } + if !g.Valid || !jsonEqual(t, g.GeoJSON, `{"type":"Point","coordinates":[1,2]}`) { + t.Errorf("geography scan = %+v", g) + } +} diff --git a/pkg/spectypes/sql_types_test.go b/pkg/spectypes/sql_types_test.go index 7e743c3..0dbf757 100644 --- a/pkg/spectypes/sql_types_test.go +++ b/pkg/spectypes/sql_types_test.go @@ -955,4 +955,3 @@ func TestSqlByteArray_Base64_RoundTrip(t *testing.T) { t.Errorf("Round-trip failed: expected %v, got %v", original, b3.Val) } } - diff --git a/pkg/spectypes/sql_vector_types.go b/pkg/spectypes/sql_vector_types.go new file mode 100644 index 0000000..bcdd6ac --- /dev/null +++ b/pkg/spectypes/sql_vector_types.go @@ -0,0 +1,308 @@ +package spectypes + +import ( + "database/sql/driver" + "encoding/json" + "fmt" + "strconv" + "strings" +) + +// pgvector column types beyond the plain `vector` (SqlVector, in +// sql_array_types.go): `halfvec`, `sparsevec` and `bit`. + +// parseVectorLiteral parses a pgvector dense literal `[1,2,3]` into []float32. +func parseVectorLiteral(s string) ([]float32, error) { + s = strings.TrimSpace(s) + if !strings.HasPrefix(s, "[") || !strings.HasSuffix(s, "]") { + return nil, fmt.Errorf("not a valid vector literal: %q", s) + } + inner := strings.TrimSpace(s[1 : len(s)-1]) + if inner == "" { + return []float32{}, nil + } + parts := strings.Split(inner, ",") + out := make([]float32, len(parts)) + for i, p := range parts { + f, err := strconv.ParseFloat(strings.TrimSpace(p), 32) + if err != nil { + return nil, fmt.Errorf("vector element %d %q: %w", i, p, err) + } + out[i] = float32(f) + } + return out, nil +} + +func formatVectorLiteral(vals []float32) string { + parts := make([]string, len(vals)) + for i, v := range vals { + parts[i] = strconv.FormatFloat(float64(v), 'f', -1, 32) + } + return "[" + strings.Join(parts, ",") + "]" +} + +// ── SqlHalfVector ──────────────────────────────────────────────────────────── + +// SqlHalfVector is a nullable pgvector `halfvec` (half-precision) column, backed +// by []float32. Wire format matches `vector`: `[1,2,3]`. +type SqlHalfVector struct { + Val []float32 + Valid bool +} + +func (v *SqlHalfVector) Scan(value any) error { + if value == nil { + v.Valid = false + v.Val = nil + return nil + } + var s string + switch val := value.(type) { + case string: + s = val + case []byte: + s = string(val) + default: + return fmt.Errorf("SqlHalfVector: cannot scan type %T", value) + } + parsed, err := parseVectorLiteral(s) + if err != nil { + return fmt.Errorf("SqlHalfVector: %w", err) + } + v.Val = parsed + v.Valid = true + return nil +} + +func (v SqlHalfVector) Value() (driver.Value, error) { + if !v.Valid { + return nil, nil + } + return formatVectorLiteral(v.Val), nil +} + +func (v SqlHalfVector) MarshalJSON() ([]byte, error) { + if !v.Valid { + return []byte("null"), nil + } + return json.Marshal(v.Val) +} + +func (v *SqlHalfVector) UnmarshalJSON(b []byte) error { + if strings.TrimSpace(string(b)) == "null" { + v.Valid = false + v.Val = nil + return nil + } + var vals []float32 + if err := json.Unmarshal(b, &vals); err != nil { + return err + } + v.Val = vals + v.Valid = true + return nil +} + +func NewSqlHalfVector(val []float32) SqlHalfVector { + return SqlHalfVector{Val: val, Valid: true} +} + +// ── SqlSparseVector ────────────────────────────────────────────────────────── + +// SqlSparseVector is a nullable pgvector `sparsevec` column. Wire format: +// `{1:0.5,4:0.2}/8` (1-based indices). JSON: +// `{"dim":8,"indices":[1,4],"values":[0.5,0.2]}`. +type SqlSparseVector struct { + Dim int + Indices []int32 + Values []float32 + Valid bool +} + +func (v *SqlSparseVector) Scan(value any) error { + if value == nil { + v.Valid = false + v.Dim, v.Indices, v.Values = 0, nil, nil + return nil + } + var s string + switch val := value.(type) { + case string: + s = val + case []byte: + s = string(val) + default: + return fmt.Errorf("SqlSparseVector: cannot scan type %T", value) + } + s = strings.TrimSpace(s) + slash := strings.LastIndex(s, "/") + if !strings.HasPrefix(s, "{") || slash < 0 || !strings.Contains(s[:slash], "}") { + return fmt.Errorf("SqlSparseVector: invalid literal %q", s) + } + dim, err := strconv.Atoi(strings.TrimSpace(s[slash+1:])) + if err != nil { + return fmt.Errorf("SqlSparseVector: bad dimension: %w", err) + } + body := strings.TrimSpace(s[1:strings.LastIndex(s, "}")]) + var idx []int32 + var vals []float32 + if body != "" { + for _, pair := range strings.Split(body, ",") { + kv := strings.SplitN(pair, ":", 2) + if len(kv) != 2 { + return fmt.Errorf("SqlSparseVector: bad pair %q", pair) + } + k, err := strconv.Atoi(strings.TrimSpace(kv[0])) + if err != nil { + return fmt.Errorf("SqlSparseVector: bad index %q: %w", kv[0], err) + } + f, err := strconv.ParseFloat(strings.TrimSpace(kv[1]), 32) + if err != nil { + return fmt.Errorf("SqlSparseVector: bad value %q: %w", kv[1], err) + } + idx = append(idx, int32(k)) + vals = append(vals, float32(f)) + } + } + v.Dim, v.Indices, v.Values, v.Valid = dim, idx, vals, true + return nil +} + +func (v SqlSparseVector) Value() (driver.Value, error) { + if !v.Valid { + return nil, nil + } + pairs := make([]string, len(v.Indices)) + for i, k := range v.Indices { + val := float32(0) + if i < len(v.Values) { + val = v.Values[i] + } + pairs[i] = strconv.Itoa(int(k)) + ":" + strconv.FormatFloat(float64(val), 'f', -1, 32) + } + return "{" + strings.Join(pairs, ",") + "}/" + strconv.Itoa(v.Dim), nil +} + +type sparseVectorJSON struct { + Dim int `json:"dim"` + Indices []int32 `json:"indices"` + Values []float32 `json:"values"` +} + +func (v SqlSparseVector) MarshalJSON() ([]byte, error) { + if !v.Valid { + return []byte("null"), nil + } + return json.Marshal(sparseVectorJSON{Dim: v.Dim, Indices: v.Indices, Values: v.Values}) +} + +func (v *SqlSparseVector) UnmarshalJSON(b []byte) error { + if strings.TrimSpace(string(b)) == "null" { + v.Valid = false + v.Dim, v.Indices, v.Values = 0, nil, nil + return nil + } + var j sparseVectorJSON + if err := json.Unmarshal(b, &j); err != nil { + return err + } + v.Dim, v.Indices, v.Values, v.Valid = j.Dim, j.Indices, j.Values, true + return nil +} + +func NewSqlSparseVector(dim int, indices []int32, values []float32) SqlSparseVector { + return SqlSparseVector{Dim: dim, Indices: indices, Values: values, Valid: true} +} + +// ── SqlBitVector ───────────────────────────────────────────────────────────── + +// SqlBitVector is a nullable Postgres `bit(n)` / `varbit` column (used by +// pgvector for Hamming/Jaccard distance), backed by []bool. Wire format: a +// string of '0'/'1' characters. JSON: a bool array. +type SqlBitVector struct { + Val []bool + Valid bool +} + +func (v *SqlBitVector) Scan(value any) error { + if value == nil { + v.Valid = false + v.Val = nil + return nil + } + var s string + switch val := value.(type) { + case string: + s = val + case []byte: + s = string(val) + default: + return fmt.Errorf("SqlBitVector: cannot scan type %T", value) + } + s = strings.TrimSpace(s) + out := make([]bool, len(s)) + for i, c := range s { + switch c { + case '1': + out[i] = true + case '0': + out[i] = false + default: + return fmt.Errorf("SqlBitVector: invalid bit %q", string(c)) + } + } + v.Val = out + v.Valid = true + return nil +} + +func (v SqlBitVector) Value() (driver.Value, error) { + if !v.Valid { + return nil, nil + } + var b strings.Builder + b.Grow(len(v.Val)) + for _, bit := range v.Val { + if bit { + b.WriteByte('1') + } else { + b.WriteByte('0') + } + } + return b.String(), nil +} + +func (v SqlBitVector) MarshalJSON() ([]byte, error) { + if !v.Valid { + return []byte("null"), nil + } + return json.Marshal(v.Val) +} + +func (v *SqlBitVector) UnmarshalJSON(b []byte) error { + s := strings.TrimSpace(string(b)) + if s == "null" { + v.Valid = false + v.Val = nil + return nil + } + // Accept both a bool array and a "0101" string. + if strings.HasPrefix(s, "\"") { + var str string + if err := json.Unmarshal(b, &str); err != nil { + return err + } + return v.Scan(str) + } + var vals []bool + if err := json.Unmarshal(b, &vals); err != nil { + return err + } + v.Val = vals + v.Valid = true + return nil +} + +func NewSqlBitVector(val []bool) SqlBitVector { + return SqlBitVector{Val: val, Valid: true} +} diff --git a/pkg/spectypes/sql_vector_types_test.go b/pkg/spectypes/sql_vector_types_test.go new file mode 100644 index 0000000..2345ee5 --- /dev/null +++ b/pkg/spectypes/sql_vector_types_test.go @@ -0,0 +1,149 @@ +package spectypes + +import ( + "encoding/json" + "reflect" + "testing" +) + +func TestSqlHalfVector_RoundTrip(t *testing.T) { + v := NewSqlHalfVector([]float32{1, 2.5, -3}) + + dv, err := v.Value() + if err != nil { + t.Fatalf("Value: %v", err) + } + if dv != "[1,2.5,-3]" { + t.Errorf("Value = %v, want [1,2.5,-3]", dv) + } + + var back SqlHalfVector + if err := back.Scan(dv.(string)); err != nil { + t.Fatalf("Scan: %v", err) + } + if !back.Valid || !reflect.DeepEqual(back.Val, v.Val) { + t.Errorf("Scan = %+v, want %+v", back, v) + } + + b, err := json.Marshal(v) + if err != nil { + t.Fatalf("Marshal: %v", err) + } + if string(b) != "[1,2.5,-3]" { + t.Errorf("json = %s, want [1,2.5,-3]", b) + } + + var fromJSON SqlHalfVector + if err := json.Unmarshal(b, &fromJSON); err != nil { + t.Fatalf("Unmarshal: %v", err) + } + if !reflect.DeepEqual(fromJSON.Val, v.Val) { + t.Errorf("json round-trip = %+v", fromJSON) + } +} + +func TestSqlHalfVector_Null(t *testing.T) { + var v SqlHalfVector + if err := v.Scan(nil); err != nil { + t.Fatalf("Scan(nil): %v", err) + } + if v.Valid { + t.Error("expected invalid after Scan(nil)") + } + dv, err := v.Value() + if err != nil || dv != nil { + t.Errorf("Value = %v, %v; want nil, nil", dv, err) + } + b, _ := json.Marshal(v) + if string(b) != "null" { + t.Errorf("json = %s, want null", b) + } +} + +func TestSqlSparseVector_RoundTrip(t *testing.T) { + v := NewSqlSparseVector(8, []int32{1, 4}, []float32{0.5, 0.2}) + + dv, err := v.Value() + if err != nil { + t.Fatalf("Value: %v", err) + } + if dv != "{1:0.5,4:0.2}/8" { + t.Errorf("Value = %v, want {1:0.5,4:0.2}/8", dv) + } + + var back SqlSparseVector + if err := back.Scan(dv.(string)); err != nil { + t.Fatalf("Scan: %v", err) + } + if back.Dim != 8 || !reflect.DeepEqual(back.Indices, []int32{1, 4}) || + !reflect.DeepEqual(back.Values, []float32{0.5, 0.2}) { + t.Errorf("Scan = %+v", back) + } + + b, err := json.Marshal(v) + if err != nil { + t.Fatalf("Marshal: %v", err) + } + want := `{"dim":8,"indices":[1,4],"values":[0.5,0.2]}` + if string(b) != want { + t.Errorf("json = %s, want %s", b, want) + } + + var fromJSON SqlSparseVector + if err := json.Unmarshal([]byte(want), &fromJSON); err != nil { + t.Fatalf("Unmarshal: %v", err) + } + if fromJSON.Dim != 8 || !fromJSON.Valid { + t.Errorf("json round-trip = %+v", fromJSON) + } +} + +func TestSqlSparseVector_ScanInvalid(t *testing.T) { + var v SqlSparseVector + for _, s := range []string{"[1,2,3]", "{1:0.5}", "{1:0.5}/x", "bad"} { + if err := v.Scan(s); err == nil { + t.Errorf("Scan(%q) expected error", s) + } + } +} + +func TestSqlBitVector_RoundTrip(t *testing.T) { + v := NewSqlBitVector([]bool{true, false, true, true}) + + dv, err := v.Value() + if err != nil { + t.Fatalf("Value: %v", err) + } + if dv != "1011" { + t.Errorf("Value = %v, want 1011", dv) + } + + var back SqlBitVector + if err := back.Scan("1011"); err != nil { + t.Fatalf("Scan: %v", err) + } + if !reflect.DeepEqual(back.Val, v.Val) { + t.Errorf("Scan = %+v", back) + } + + b, _ := json.Marshal(v) + if string(b) != "[true,false,true,true]" { + t.Errorf("json = %s", b) + } + + // JSON also accepts a "0101" string. + var fromStr SqlBitVector + if err := json.Unmarshal([]byte(`"1011"`), &fromStr); err != nil { + t.Fatalf("Unmarshal string: %v", err) + } + if !reflect.DeepEqual(fromStr.Val, v.Val) { + t.Errorf("string json = %+v", fromStr) + } +} + +func TestSqlBitVector_ScanInvalid(t *testing.T) { + var v SqlBitVector + if err := v.Scan("1021"); err == nil { + t.Error("expected error for invalid bit") + } +} diff --git a/pkg/spectypes/type_names.go b/pkg/spectypes/type_names.go new file mode 100644 index 0000000..471224c --- /dev/null +++ b/pkg/spectypes/type_names.go @@ -0,0 +1,93 @@ +package spectypes + +import ( + "reflect" + "strings" +) + +// pkgPath is the import path of this package, used to recognise spectypes +// wrappers by reflection. +const pkgPath = "github.com/bitechdev/ResolveSpec/pkg/spectypes" + +// canonicalSQLNames maps a spectypes wrapper type name to the PostgreSQL type +// name it represents. Dimensioned types (vector(1536), geometry(Point,4326)) +// still need a gorm/bun `type:` tag for the full declaration — this is the +// fallback used for metadata and OpenAPI when no tag is present. +var canonicalSQLNames = map[string]string{ + "SqlVector": "vector", + "SqlHalfVector": "halfvec", + "SqlSparseVector": "sparsevec", + "SqlBitVector": "bit", + "SqlGeometry": "geometry", + "SqlGeography": "geography", + "SqlJSONB": "jsonb", + "SqlStringArray": "text[]", + "SqlInt16Array": "smallint[]", + "SqlInt32Array": "integer[]", + "SqlInt64Array": "bigint[]", + "SqlFloat32Array": "real[]", + "SqlFloat64Array": "double precision[]", + "SqlBoolArray": "boolean[]", + "SqlUUIDArray": "uuid[]", + "SqlDate": "date", + "SqlTime": "time", + "SqlTimeStamp": "timestamp", +} + +// sqlNullElemNames maps the element type of a SqlNull[T] alias to a PG type name. +var sqlNullElemNames = map[string]string{ + "int16": "smallint", + "int32": "integer", + "int64": "bigint", + "float64": "double precision", + "bool": "boolean", + "string": "text", + "[]uint8": "bytea", + "uuid.UUID": "uuid", + "Time": "timestamp", +} + +// SQLTypeName returns the canonical PostgreSQL type name for a spectypes wrapper +// type, or ("", false) if t is not a recognised spectypes type. +func SQLTypeName(t reflect.Type) (string, bool) { + for t != nil && t.Kind() == reflect.Pointer { + t = t.Elem() + } + if t == nil || t.PkgPath() != pkgPath { + return "", false + } + + name := t.Name() + if n, ok := canonicalSQLNames[name]; ok { + return n, true + } + + // SqlNull[T] aliases, e.g. "SqlNull[int16]", "SqlNull[uuid.UUID]". + if strings.HasPrefix(name, "SqlNull[") && strings.HasSuffix(name, "]") { + elem := name[len("SqlNull[") : len(name)-1] + if idx := strings.LastIndex(elem, "."); idx >= 0 { + // keep last path segment, e.g. "github.com/google/uuid.UUID" -> "uuid.UUID" + if slash := strings.LastIndex(elem[:idx], "/"); slash >= 0 { + elem = elem[slash+1:] + } + } + if n, ok := sqlNullElemNames[elem]; ok { + return n, true + } + return "", false + } + + return "", false +} + +// IsSpatialType reports whether t is a PostGIS geometry/geography wrapper. +func IsSpatialType(t reflect.Type) bool { + n, ok := SQLTypeName(t) + return ok && (n == "geometry" || n == "geography") +} + +// IsVectorType reports whether t is a pgvector wrapper (vector/halfvec/sparsevec). +func IsVectorType(t reflect.Type) bool { + n, ok := SQLTypeName(t) + return ok && (n == "vector" || n == "halfvec" || n == "sparsevec") +} diff --git a/pkg/spectypes/type_names_test.go b/pkg/spectypes/type_names_test.go new file mode 100644 index 0000000..fdda16f --- /dev/null +++ b/pkg/spectypes/type_names_test.go @@ -0,0 +1,66 @@ +package spectypes + +import ( + "reflect" + "testing" +) + +func TestSQLTypeName(t *testing.T) { + cases := []struct { + val any + want string + }{ + {SqlVector{}, "vector"}, + {SqlHalfVector{}, "halfvec"}, + {SqlSparseVector{}, "sparsevec"}, + {SqlBitVector{}, "bit"}, + {SqlGeometry{}, "geometry"}, + {SqlGeography{}, "geography"}, + {SqlJSONB{}, "jsonb"}, + {SqlStringArray{}, "text[]"}, + {SqlString{}, "text"}, + {SqlInt64{}, "bigint"}, + } + for _, c := range cases { + got, ok := SQLTypeName(reflect.TypeOf(c.val)) + if !ok || got != c.want { + t.Errorf("SQLTypeName(%T) = %q, %v; want %q", c.val, got, ok, c.want) + } + } + + // Pointer is unwrapped. + if got, ok := SQLTypeName(reflect.TypeOf(&SqlGeometry{})); !ok || got != "geometry" { + t.Errorf("pointer: got %q, %v", got, ok) + } + + // Non-spectypes type. + if _, ok := SQLTypeName(reflect.TypeOf("")); ok { + t.Error("expected false for string") + } +} + +func TestIsSpatialType(t *testing.T) { + if !IsSpatialType(reflect.TypeOf(SqlGeometry{})) { + t.Error("SqlGeometry should be spatial") + } + if !IsSpatialType(reflect.TypeOf(SqlGeography{})) { + t.Error("SqlGeography should be spatial") + } + if IsSpatialType(reflect.TypeOf(SqlVector{})) { + t.Error("SqlVector should not be spatial") + } +} + +func TestIsVectorType(t *testing.T) { + for _, v := range []any{SqlVector{}, SqlHalfVector{}, SqlSparseVector{}} { + if !IsVectorType(reflect.TypeOf(v)) { + t.Errorf("%T should be vector", v) + } + } + if IsVectorType(reflect.TypeOf(SqlBitVector{})) { + t.Error("SqlBitVector is not a vector type") + } + if IsVectorType(reflect.TypeOf(SqlGeometry{})) { + t.Error("SqlGeometry is not a vector type") + } +} diff --git a/pkg/spectypes/wkb.go b/pkg/spectypes/wkb.go new file mode 100644 index 0000000..92f0dde --- /dev/null +++ b/pkg/spectypes/wkb.go @@ -0,0 +1,697 @@ +package spectypes + +import ( + "encoding/binary" + "encoding/hex" + "encoding/json" + "fmt" + "math" + "strconv" + "strings" +) + +// Minimal self-contained EWKB (PostGIS extended WKB) <-> GeoJSON / WKT codec. +// Supports 2D and 3D (Z) geometries of type Point, LineString, Polygon, +// MultiPoint, MultiLineString, MultiPolygon and GeometryCollection. The M +// dimension is parsed but dropped (GeoJSON has no M). SRID is tracked separately +// from the GeoJSON payload (GeoJSON assumes CRS84 / EPSG:4326). + +// EWKB type flag bits (PostGIS). +const ( + ewkbZ = 0x80000000 + ewkbM = 0x40000000 + ewkbSRID = 0x20000000 +) + +// geom is the intermediate geometry representation used by the codec. +// +// Point -> coord ([]float64, len 2 or 3) +// LineString/MultiPt -> line ([][]float64) +// Polygon/MultiLine -> poly ([][][]float64) +// MultiPolygon -> multi ([][][][]float64) +// GeometryCollection -> geoms ([]geom) +type geom struct { + typ string + coord []float64 + line [][]float64 + poly [][][]float64 + multi [][][][]float64 + geoms []geom +} + +// wkbReader consumes an EWKB byte stream. +type wkbReader struct { + buf []byte + pos int +} + +func (r *wkbReader) readByte() (byte, error) { + if r.pos >= len(r.buf) { + return 0, fmt.Errorf("wkb: unexpected end of input") + } + b := r.buf[r.pos] + r.pos++ + return b, nil +} + +func (r *wkbReader) readUint32(bo binary.ByteOrder) (uint32, error) { + if r.pos+4 > len(r.buf) { + return 0, fmt.Errorf("wkb: unexpected end of input") + } + v := bo.Uint32(r.buf[r.pos:]) + r.pos += 4 + return v, nil +} + +func (r *wkbReader) readFloat64(bo binary.ByteOrder) (float64, error) { + if r.pos+8 > len(r.buf) { + return 0, fmt.Errorf("wkb: unexpected end of input") + } + v := math.Float64frombits(bo.Uint64(r.buf[r.pos:])) + r.pos += 8 + return v, nil +} + +// DecodeEWKBHex decodes a PostGIS hex-EWKB string (the default text +// representation of a geometry column) into a GeoJSON geometry object and its +// SRID. An SRID of 0 means "unspecified". +func DecodeEWKBHex(s string) (geojson []byte, srid int, err error) { + s = strings.TrimSpace(s) + raw, err := hex.DecodeString(s) + if err != nil { + return nil, 0, fmt.Errorf("wkb: invalid hex: %w", err) + } + return DecodeEWKB(raw) +} + +// DecodeEWKB decodes raw PostGIS EWKB bytes into a GeoJSON geometry object and +// its SRID. +func DecodeEWKB(raw []byte) (geojson []byte, srid int, err error) { + r := &wkbReader{buf: raw} + g, sr, err := readGeom(r) + if err != nil { + return nil, 0, err + } + out, err := geomToGeoJSON(g) + if err != nil { + return nil, 0, err + } + return out, sr, nil +} + +func readGeom(r *wkbReader) (geom, int, error) { + order, err := r.readByte() + if err != nil { + return geom{}, 0, err + } + var bo binary.ByteOrder + switch order { + case 0: + bo = binary.BigEndian + case 1: + bo = binary.LittleEndian + default: + return geom{}, 0, fmt.Errorf("wkb: invalid byte order %d", order) + } + + rawType, err := r.readUint32(bo) + if err != nil { + return geom{}, 0, err + } + hasZ := rawType&ewkbZ != 0 + hasM := rawType&ewkbM != 0 + hasSRID := rawType&ewkbSRID != 0 + baseType := rawType & 0xff + + srid := 0 + if hasSRID { + s, err := r.readUint32(bo) + if err != nil { + return geom{}, 0, err + } + srid = int(s) + } + + dims := 2 + if hasZ { + dims = 3 + } + // M is consumed but not retained. + stride := dims + if hasM { + stride++ + } + + readCoord := func() ([]float64, error) { + c := make([]float64, 0, dims) + for i := 0; i < stride; i++ { + v, err := r.readFloat64(bo) + if err != nil { + return nil, err + } + if i < dims { + c = append(c, v) + } + } + return c, nil + } + readLine := func() ([][]float64, error) { + n, err := r.readUint32(bo) + if err != nil { + return nil, err + } + pts := make([][]float64, n) + for i := range pts { + pts[i], err = readCoord() + if err != nil { + return nil, err + } + } + return pts, nil + } + readPoly := func() ([][][]float64, error) { + n, err := r.readUint32(bo) + if err != nil { + return nil, err + } + rings := make([][][]float64, n) + for i := range rings { + rings[i], err = readLine() + if err != nil { + return nil, err + } + } + return rings, nil + } + + switch baseType { + case 1: // Point + c, err := readCoord() + if err != nil { + return geom{}, 0, err + } + return geom{typ: "Point", coord: c}, srid, nil + case 2: // LineString + l, err := readLine() + if err != nil { + return geom{}, 0, err + } + return geom{typ: "LineString", line: l}, srid, nil + case 3: // Polygon + p, err := readPoly() + if err != nil { + return geom{}, 0, err + } + return geom{typ: "Polygon", poly: p}, srid, nil + case 4, 5, 6: // Multi* + n, err := r.readUint32(bo) + if err != nil { + return geom{}, 0, err + } + parts := make([]geom, n) + for i := range parts { + sub, _, err := readGeom(r) + if err != nil { + return geom{}, 0, err + } + parts[i] = sub + } + switch baseType { + case 4: + pts := make([][]float64, len(parts)) + for i, p := range parts { + pts[i] = p.coord + } + return geom{typ: "MultiPoint", line: pts}, srid, nil + case 5: + lines := make([][][]float64, len(parts)) + for i, p := range parts { + lines[i] = p.line + } + return geom{typ: "MultiLineString", poly: lines}, srid, nil + default: + polys := make([][][][]float64, len(parts)) + for i, p := range parts { + polys[i] = p.poly + } + return geom{typ: "MultiPolygon", multi: polys}, srid, nil + } + case 7: // GeometryCollection + n, err := r.readUint32(bo) + if err != nil { + return geom{}, 0, err + } + parts := make([]geom, n) + for i := range parts { + sub, _, err := readGeom(r) + if err != nil { + return geom{}, 0, err + } + parts[i] = sub + } + return geom{typ: "GeometryCollection", geoms: parts}, srid, nil + default: + return geom{}, 0, fmt.Errorf("wkb: unsupported geometry type %d", baseType) + } +} + +// ── GeoJSON ────────────────────────────────────────────────────────────────── + +type geoJSON struct { + Type string `json:"type"` + Coordinates json.RawMessage `json:"coordinates,omitempty"` + Geometries []geoJSON `json:"geometries,omitempty"` +} + +func geomToGeoJSON(g geom) ([]byte, error) { + gj, err := geomToGeoJSONStruct(g) + if err != nil { + return nil, err + } + return json.Marshal(gj) +} + +func geomToGeoJSONStruct(g geom) (geoJSON, error) { + var coords any + switch g.typ { + case "Point": + coords = g.coord + case "LineString", "MultiPoint": + coords = g.line + case "Polygon", "MultiLineString": + coords = g.poly + case "MultiPolygon": + coords = g.multi + case "GeometryCollection": + subs := make([]geoJSON, len(g.geoms)) + for i, sub := range g.geoms { + s, err := geomToGeoJSONStruct(sub) + if err != nil { + return geoJSON{}, err + } + subs[i] = s + } + return geoJSON{Type: "GeometryCollection", Geometries: subs}, nil + default: + return geoJSON{}, fmt.Errorf("wkb: cannot encode geometry type %q", g.typ) + } + rc, err := json.Marshal(coords) + if err != nil { + return geoJSON{}, err + } + return geoJSON{Type: g.typ, Coordinates: rc}, nil +} + +func geoJSONToGeom(data []byte) (geom, error) { + var gj geoJSON + if err := json.Unmarshal(data, &gj); err != nil { + return geom{}, fmt.Errorf("geojson: %w", err) + } + return geoJSONStructToGeom(gj) +} + +func geoJSONStructToGeom(gj geoJSON) (geom, error) { + switch gj.Type { + case "Point": + var c []float64 + if err := json.Unmarshal(gj.Coordinates, &c); err != nil { + return geom{}, err + } + return geom{typ: "Point", coord: c}, nil + case "LineString", "MultiPoint": + var l [][]float64 + if err := json.Unmarshal(gj.Coordinates, &l); err != nil { + return geom{}, err + } + return geom{typ: gj.Type, line: l}, nil + case "Polygon", "MultiLineString": + var p [][][]float64 + if err := json.Unmarshal(gj.Coordinates, &p); err != nil { + return geom{}, err + } + return geom{typ: gj.Type, poly: p}, nil + case "MultiPolygon": + var m [][][][]float64 + if err := json.Unmarshal(gj.Coordinates, &m); err != nil { + return geom{}, err + } + return geom{typ: gj.Type, multi: m}, nil + case "GeometryCollection": + subs := make([]geom, len(gj.Geometries)) + for i, s := range gj.Geometries { + g, err := geoJSONStructToGeom(s) + if err != nil { + return geom{}, err + } + subs[i] = g + } + return geom{typ: "GeometryCollection", geoms: subs}, nil + default: + return geom{}, fmt.Errorf("geojson: unsupported type %q", gj.Type) + } +} + +// ── WKT ────────────────────────────────────────────────────────────────────── + +// GeoJSONToWKT converts a GeoJSON geometry object to its WKT representation. +func GeoJSONToWKT(geojson []byte) (string, error) { + g, err := geoJSONToGeom(geojson) + if err != nil { + return "", err + } + return geomToWKT(g) +} + +func fmtNum(f float64) string { + return strconv.FormatFloat(f, 'f', -1, 64) +} + +func coordWKT(c []float64) string { + parts := make([]string, len(c)) + for i, v := range c { + parts[i] = fmtNum(v) + } + return strings.Join(parts, " ") +} + +func lineWKT(pts [][]float64) string { + parts := make([]string, len(pts)) + for i, p := range pts { + parts[i] = coordWKT(p) + } + return "(" + strings.Join(parts, ", ") + ")" +} + +func polyWKT(rings [][][]float64) string { + parts := make([]string, len(rings)) + for i, r := range rings { + parts[i] = lineWKT(r) + } + return "(" + strings.Join(parts, ", ") + ")" +} + +// ── WKT parsing ───────────────────────────────────────────────────────────── + +// parseWKT parses a subset of WKT (2D/3D, no M) into the intermediate geom. +func parseWKT(s string) (geom, error) { + p := &wktParser{s: s} + p.skipSpace() + g, err := p.parseGeom() + if err != nil { + return geom{}, err + } + p.skipSpace() + if p.pos != len(p.s) { + return geom{}, fmt.Errorf("wkt: trailing input %q", p.s[p.pos:]) + } + return g, nil +} + +type wktParser struct { + s string + pos int +} + +func (p *wktParser) skipSpace() { + for p.pos < len(p.s) && (p.s[p.pos] == ' ' || p.s[p.pos] == '\t' || p.s[p.pos] == '\n' || p.s[p.pos] == '\r') { + p.pos++ + } +} + +func (p *wktParser) parseGeom() (geom, error) { + p.skipSpace() + start := p.pos + for p.pos < len(p.s) && (p.s[p.pos] >= 'A' && p.s[p.pos] <= 'Z' || p.s[p.pos] >= 'a' && p.s[p.pos] <= 'z') { + p.pos++ + } + kw := strings.ToUpper(p.s[start:p.pos]) + p.skipSpace() + // Optional Z / M / ZM dimension tag — coordinates carry their own arity. + if p.pos < len(p.s) && (p.s[p.pos] == 'Z' || p.s[p.pos] == 'M' || p.s[p.pos] == 'z' || p.s[p.pos] == 'm') { + for p.pos < len(p.s) && p.s[p.pos] != '(' { + p.pos++ + } + } + p.skipSpace() + + switch kw { + case "POINT": + pts, err := p.parseCoordList() + if err != nil { + return geom{}, err + } + if len(pts) != 1 { + return geom{}, fmt.Errorf("wkt: POINT needs exactly one coordinate") + } + return geom{typ: "Point", coord: pts[0]}, nil + case "LINESTRING": + pts, err := p.parseCoordList() + if err != nil { + return geom{}, err + } + return geom{typ: "LineString", line: pts}, nil + case "MULTIPOINT": + pts, err := p.parseMultiPoint() + if err != nil { + return geom{}, err + } + return geom{typ: "MultiPoint", line: pts}, nil + case "POLYGON": + rings, err := p.parseRingList() + if err != nil { + return geom{}, err + } + return geom{typ: "Polygon", poly: rings}, nil + case "MULTILINESTRING": + lines, err := p.parseRingList() + if err != nil { + return geom{}, err + } + return geom{typ: "MultiLineString", poly: lines}, nil + case "MULTIPOLYGON": + polys, err := p.parsePolyList() + if err != nil { + return geom{}, err + } + return geom{typ: "MultiPolygon", multi: polys}, nil + case "GEOMETRYCOLLECTION": + return p.parseCollection() + default: + return geom{}, fmt.Errorf("wkt: unsupported geometry %q", kw) + } +} + +func (p *wktParser) expect(c byte) error { + p.skipSpace() + if p.pos >= len(p.s) || p.s[p.pos] != c { + return fmt.Errorf("wkt: expected %q at offset %d", string(c), p.pos) + } + p.pos++ + return nil +} + +func (p *wktParser) peek() byte { + p.skipSpace() + if p.pos >= len(p.s) { + return 0 + } + return p.s[p.pos] +} + +// parseCoordList parses `(x y[, x y]...)`. +func (p *wktParser) parseCoordList() ([][]float64, error) { + if err := p.expect('('); err != nil { + return nil, err + } + var out [][]float64 + for { + c, err := p.parseCoord() + if err != nil { + return nil, err + } + out = append(out, c) + if p.peek() == ',' { + p.pos++ + continue + } + break + } + if err := p.expect(')'); err != nil { + return nil, err + } + return out, nil +} + +func (p *wktParser) parseCoord() ([]float64, error) { + p.skipSpace() + // Some MULTIPOINT forms wrap each coord in parentheses. + wrapped := false + if p.peek() == '(' { + p.pos++ + wrapped = true + } + var nums []float64 + for { + p.skipSpace() + start := p.pos + for p.pos < len(p.s) { + ch := p.s[p.pos] + if ch == '-' || ch == '+' || ch == '.' || ch == 'e' || ch == 'E' || (ch >= '0' && ch <= '9') { + p.pos++ + continue + } + break + } + if p.pos == start { + break + } + f, err := strconv.ParseFloat(p.s[start:p.pos], 64) + if err != nil { + return nil, fmt.Errorf("wkt: bad number %q", p.s[start:p.pos]) + } + nums = append(nums, f) + p.skipSpace() + if p.pos < len(p.s) && p.s[p.pos] == ' ' { + continue + } + } + if wrapped { + if err := p.expect(')'); err != nil { + return nil, err + } + } + if len(nums) < 2 { + return nil, fmt.Errorf("wkt: coordinate needs at least 2 numbers") + } + if len(nums) > 3 { + nums = nums[:3] + } + return nums, nil +} + +func (p *wktParser) parseMultiPoint() ([][]float64, error) { + if err := p.expect('('); err != nil { + return nil, err + } + var out [][]float64 + for { + c, err := p.parseCoord() + if err != nil { + return nil, err + } + out = append(out, c) + if p.peek() == ',' { + p.pos++ + continue + } + break + } + if err := p.expect(')'); err != nil { + return nil, err + } + return out, nil +} + +// parseRingList parses `((x y, ...), (...))`. +func (p *wktParser) parseRingList() ([][][]float64, error) { + if err := p.expect('('); err != nil { + return nil, err + } + var out [][][]float64 + for { + ring, err := p.parseCoordList() + if err != nil { + return nil, err + } + out = append(out, ring) + if p.peek() == ',' { + p.pos++ + continue + } + break + } + if err := p.expect(')'); err != nil { + return nil, err + } + return out, nil +} + +// parsePolyList parses `(((...)), ((...)))`. +func (p *wktParser) parsePolyList() ([][][][]float64, error) { + if err := p.expect('('); err != nil { + return nil, err + } + var out [][][][]float64 + for { + poly, err := p.parseRingList() + if err != nil { + return nil, err + } + out = append(out, poly) + if p.peek() == ',' { + p.pos++ + continue + } + break + } + if err := p.expect(')'); err != nil { + return nil, err + } + return out, nil +} + +func (p *wktParser) parseCollection() (geom, error) { + if err := p.expect('('); err != nil { + return geom{}, err + } + var subs []geom + for { + g, err := p.parseGeom() + if err != nil { + return geom{}, err + } + subs = append(subs, g) + if p.peek() == ',' { + p.pos++ + continue + } + break + } + if err := p.expect(')'); err != nil { + return geom{}, err + } + return geom{typ: "GeometryCollection", geoms: subs}, nil +} + +func geomToWKT(g geom) (string, error) { + switch g.typ { + case "Point": + return "POINT (" + coordWKT(g.coord) + ")", nil + case "LineString": + return "LINESTRING " + lineWKT(g.line), nil + case "MultiPoint": + return "MULTIPOINT " + lineWKT(g.line), nil + case "Polygon": + return "POLYGON " + polyWKT(g.poly), nil + case "MultiLineString": + return "MULTILINESTRING " + polyWKT(g.poly), nil + case "MultiPolygon": + parts := make([]string, len(g.multi)) + for i, p := range g.multi { + parts[i] = polyWKT(p) + } + return "MULTIPOLYGON (" + strings.Join(parts, ", ") + ")", nil + case "GeometryCollection": + parts := make([]string, len(g.geoms)) + for i, sub := range g.geoms { + w, err := geomToWKT(sub) + if err != nil { + return "", err + } + parts[i] = w + } + return "GEOMETRYCOLLECTION (" + strings.Join(parts, ", ") + ")", nil + default: + return "", fmt.Errorf("wkt: cannot encode geometry type %q", g.typ) + } +} diff --git a/pkg/spectypes/wkb_test.go b/pkg/spectypes/wkb_test.go new file mode 100644 index 0000000..336b76a --- /dev/null +++ b/pkg/spectypes/wkb_test.go @@ -0,0 +1,91 @@ +package spectypes + +import ( + "encoding/json" + "testing" +) + +func TestDecodeEWKBHex(t *testing.T) { + tests := []struct { + name string + hex string + wantSRID int + wantJSON string + }{ + { + // SRID=4326;POINT(1 2) + name: "point with srid", + hex: "0101000020E6100000000000000000F03F0000000000000040", + wantSRID: 4326, + wantJSON: `{"type":"Point","coordinates":[1,2]}`, + }, + { + // POINT(1 2) no SRID, little endian + name: "point no srid", + hex: "0101000000000000000000F03F0000000000000040", + wantSRID: 0, + wantJSON: `{"type":"Point","coordinates":[1,2]}`, + }, + { + // SRID=4326;LINESTRING(0 0, 1 1, 2 2) + name: "linestring", + hex: "0102000020E610000003000000000000000000000000000000000000000000000000" + + "00F03F000000000000F03F00000000000000400000000000000040", + wantSRID: 4326, + wantJSON: `{"type":"LineString","coordinates":[[0,0],[1,1],[2,2]]}`, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + gj, srid, err := DecodeEWKBHex(tt.hex) + if err != nil { + t.Fatalf("DecodeEWKBHex: %v", err) + } + if srid != tt.wantSRID { + t.Errorf("srid = %d, want %d", srid, tt.wantSRID) + } + if !jsonEqual(t, gj, tt.wantJSON) { + t.Errorf("geojson = %s, want %s", gj, tt.wantJSON) + } + // Round-trip through WKT parser. + wkt, err := GeoJSONToWKT(gj) + if err != nil { + t.Fatalf("GeoJSONToWKT: %v", err) + } + gj2, err := wktToGeoJSON(wkt) + if err != nil { + t.Fatalf("wktToGeoJSON(%q): %v", wkt, err) + } + if !jsonEqual(t, gj2, tt.wantJSON) { + t.Errorf("round-trip geojson = %s, want %s", gj2, tt.wantJSON) + } + }) + } +} + +func TestWKTPolygon(t *testing.T) { + src := `POLYGON ((0 0, 4 0, 4 4, 0 4, 0 0), (1 1, 2 1, 2 2, 1 2, 1 1))` + gj, err := wktToGeoJSON(src) + if err != nil { + t.Fatalf("wktToGeoJSON: %v", err) + } + want := `{"type":"Polygon","coordinates":[[[0,0],[4,0],[4,4],[0,4],[0,0]],[[1,1],[2,1],[2,2],[1,2],[1,1]]]}` + if !jsonEqual(t, gj, want) { + t.Errorf("geojson = %s, want %s", gj, want) + } +} + +func jsonEqual(t *testing.T, got []byte, want string) bool { + t.Helper() + var a, b any + if err := json.Unmarshal(got, &a); err != nil { + t.Fatalf("unmarshal got %s: %v", got, err) + } + if err := json.Unmarshal([]byte(want), &b); err != nil { + t.Fatalf("unmarshal want %s: %v", want, err) + } + ab, _ := json.Marshal(a) + bb, _ := json.Marshal(b) + return string(ab) == string(bb) +} diff --git a/todo.md b/todo.md index da8fc52..d83fb3c 100644 --- a/todo.md +++ b/todo.md @@ -68,6 +68,14 @@ See [`resolvespec-python/todo.md`](./resolvespec-python/todo.md) for detailed Py - [ ] Computed column improvements - [ ] Recursive query support +3. **PostGIS & pgvector (PostgreSQL)** + - [x] Custom types: `SqlGeometry`/`SqlGeography` (GeoJSON I/O), `SqlHalfVector`, `SqlSparseVector`, `SqlBitVector` + - [x] Spatial filter operators (`st_dwithin`, `st_intersects`, `bbox`, …) in resolvespec + restheadspec + - [x] Vector similarity filter operators (`l2_within`, `cosine_within`, `ip_within`) + - [x] KNN structured option (`options.vector_search` / `X-Vector-Search-*` headers) + - [x] Metadata + OpenAPI report geometry/vector column types + - [ ] Integration tests against a PostGIS + pgvector database (needs test DB with extensions) + 3. **Testing & Quality** - [ ] Increase test coverage to 70%+ - [ ] Add integration tests for all ORMs @@ -96,5 +104,5 @@ See [`resolvespec-python/todo.md`](./resolvespec-python/todo.md) for detailed Py --- -**Last Updated:** 2026-02-07 -**Updated:** Added resolvespec-js client testing and implementation tasks +**Last Updated:** 2026-08-29 +**Updated:** Added PostGIS + pgvector support (types, filter operators, KNN, metadata/OpenAPI)