mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-08-31 05:22:35 +00:00
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
This commit is contained in:
@@ -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=<n>;<WKT>` 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=<n>;<WKT>`.
|
||||
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=<n>;` 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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -955,4 +955,3 @@ func TestSqlByteArray_Base64_RoundTrip(t *testing.T) {
|
||||
t.Errorf("Round-trip failed: expected %v, got %v", original, b3.Val)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
Reference in New Issue
Block a user