diff --git a/go.mod b/go.mod index c89cfa38f8..d8633c12c2 100644 --- a/go.mod +++ b/go.mod @@ -30,27 +30,36 @@ require ( google.golang.org/grpc v1.81.1 google.golang.org/protobuf v1.36.11 gopkg.in/yaml.v3 v3.0.1 + modernc.org/sqlite v1.48.1 ) require ( cel.dev/expr v0.25.1 // indirect filippo.io/edwards25519 v1.2.0 // indirect + github.com/dustin/go-humanize v1.0.1 // indirect + github.com/google/uuid v1.6.0 // indirect github.com/inconshreveable/mousetrap v1.1.0 // indirect github.com/jackc/pgpassfile v1.0.0 // indirect github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect github.com/jackc/puddle/v2 v2.2.2 // indirect github.com/kr/text v0.2.0 // indirect + github.com/mattn/go-isatty v0.0.20 // indirect github.com/ncruces/go-sqlite3-wasm/v2 v2.6.35302 // indirect + github.com/ncruces/go-strftime v1.0.0 // indirect github.com/ncruces/julianday v1.0.0 // indirect + github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect github.com/rogpeppe/go-internal v1.10.0 // indirect github.com/wasilibs/wazero-helpers v0.0.0-20240620070341-3dff1577cd52 // indirect github.com/xeipuuv/gojsonpointer v0.0.0-20180127040702-4e3ac2762d5f // indirect github.com/xeipuuv/gojsonreference v0.0.0-20180127040603-bd5ef7bd5415 // indirect go.yaml.in/yaml/v3 v3.0.4 // indirect - golang.org/x/exp v0.0.0-20250620022241-b7579e27df2b // indirect + golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546 // indirect golang.org/x/net v0.55.0 // indirect golang.org/x/sys v0.45.0 // indirect golang.org/x/text v0.37.0 // indirect google.golang.org/genproto/googleapis/api v0.0.0-20260226221140-a57be14db171 // indirect google.golang.org/genproto/googleapis/rpc v0.0.0-20260226221140-a57be14db171 // indirect + modernc.org/libc v1.70.0 // indirect + modernc.org/mathutil v1.7.1 // indirect + modernc.org/memory v1.11.0 // indirect ) diff --git a/go.sum b/go.sum index 017fc2647b..3298bf1540 100644 --- a/go.sum +++ b/go.sum @@ -13,6 +13,8 @@ github.com/cubicdaiya/gonp v1.0.4/go.mod h1:iWGuP/7+JVTn02OWhRemVbMmG1DOUnmrGTYY github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY= +github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto= github.com/fatih/structtag v1.2.0 h1:/OdNE99OxoI/PqaW/SuSK9uxxT3f/tcSZgon/ssNSx4= github.com/fatih/structtag v1.2.0/go.mod h1:mBJUNpUnHmRKrKlQQlmCrh5PuhftFbNv8Ys4/aAZl94= github.com/go-logr/logr v1.4.3 h1:CjnDlHq8ikf6E492q6eKboGOC0T8CDaOvkHCIg8idEI= @@ -27,8 +29,12 @@ github.com/google/cel-go v0.28.1 h1:YWIwi77J4xIsYUwAF/iIuS6haffzIHS8yWI8glSbLWM= github.com/google/cel-go v0.28.1/go.mod h1:X0bD6iVNR8pkROSOoHVdgTkzmRcosof7WQqCD6wcMc8= github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= +github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e h1:ijClszYn+mADRFY17kjQEVQ1XRhq2/JR1M3sGqeJoxs= +github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e/go.mod h1:boTsfXsheKC2y+lKOCMpSfarhxDeIzfZG1jqGcPl3cA= github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k= +github.com/hashicorp/golang-lru/v2 v2.0.7/go.mod h1:QeFd9opnmA6QUJc5vARoKUSoFhyfM2/ZepoAG6RGpeM= github.com/inconshreveable/mousetrap v1.1.0 h1:wN+x4NVGpMsO7ErUn/mUI3vEoE6Jt13X2s0bqwp9tc8= github.com/inconshreveable/mousetrap v1.1.0/go.mod h1:vpF70FUmC8bwa3OWnCshd2FqLfsEA9PFc4w1p2J65bw= github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsIM= @@ -47,16 +53,22 @@ github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE= github.com/lib/pq v1.12.3 h1:tTWxr2YLKwIvK90ZXEw8GP7UFHtcbTtty8zsI+YjrfQ= github.com/lib/pq v1.12.3/go.mod h1:/p+8NSbOcwzAEI7wiMXFlgydTwcgTr3OSKMsD2BitpA= +github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY= +github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y= github.com/ncruces/go-sqlite3 v0.34.4 h1:bp8jd1o2CMvjkrrp1lOtHDu1TdzExNWACt+piQeMwzg= github.com/ncruces/go-sqlite3 v0.34.4/go.mod h1:tOyhDWnzlrzflKIGw177imjuO4ZEbfOS1NJ67B1BanQ= github.com/ncruces/go-sqlite3-wasm/v2 v2.6.35302 h1:IZCiInPIp6OhOc1skMDGOBMwMQuE1TK+QqVe35vd/ro= github.com/ncruces/go-sqlite3-wasm/v2 v2.6.35302/go.mod h1:ELHF6yqC51E0DiitfabHXl/aKKouCihugbhNT5a+yEY= +github.com/ncruces/go-strftime v1.0.0 h1:HMFp8mLCTPp341M/ZnA4qaf7ZlsbTc+miZjCLOFAw7w= +github.com/ncruces/go-strftime v1.0.0/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls= github.com/ncruces/julianday v1.0.0 h1:fH0OKwa7NWvniGQtxdJRxAgkBMolni2BjDHaWTxqt7M= github.com/ncruces/julianday v1.0.0/go.mod h1:Dusn2KvZrrovOMJuOt0TNXL6tB7U2E8kvza5fFc9G7g= github.com/pganalyze/pg_query_go/v6 v6.2.2 h1:O0L6zMC226R82RF3X5n0Ki6HjytDsoAzuzp4ATVAHNo= github.com/pganalyze/pg_query_go/v6 v6.2.2/go.mod h1:Cn6+j4870kJz3iYNsb0VsNG04vpSWgEvBwc590J4qD0= github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= +github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE= +github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo= github.com/riza-io/grpc-go v0.2.0 h1:2HxQKFVE7VuYstcJ8zqpN84VnAoJ4dCL6YFhJewNcHQ= github.com/riza-io/grpc-go v0.2.0/go.mod h1:2bDvR9KkKC3KhtlSHfR3dAXjUMT86kg4UfWFyVGWqi8= github.com/rogpeppe/go-internal v1.10.0 h1:TMyTOH3F/DB16zRVcYyreMH6GnZZrwQVAoYjRBZyWFQ= @@ -104,16 +116,21 @@ go.opentelemetry.io/otel/trace v1.43.0 h1:BkNrHpup+4k4w+ZZ86CZoHHEkohws8AY+WTX09 go.opentelemetry.io/otel/trace v1.43.0/go.mod h1:/QJhyVBUUswCphDVxq+8mld+AvhXZLhe+8WVFxiFff0= go.yaml.in/yaml/v3 v3.0.4 h1:tfq32ie2Jv2UxXFdLJdh3jXuOzWiL1fo0bu/FbuKpbc= go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg= -golang.org/x/exp v0.0.0-20250620022241-b7579e27df2b h1:M2rDM6z3Fhozi9O7NWsxAkg/yqS/lQJ6PmkyIV3YP+o= -golang.org/x/exp v0.0.0-20250620022241-b7579e27df2b/go.mod h1:3//PLf8L/X+8b4vuAfHzxeRUl04Adcb341+IGKfnqS8= +golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546 h1:mgKeJMpvi0yx/sU5GsxQ7p6s2wtOnGAHZWCHUM4KGzY= +golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546/go.mod h1:j/pmGrbnkbPtQfxEe5D0VQhZC6qKbfKifgD0oM7sR70= +golang.org/x/mod v0.35.0 h1:Ww1D637e6Pg+Zb2KrWfHQUnH2dQRLBQyAtpr/haaJeM= +golang.org/x/mod v0.35.0/go.mod h1:+GwiRhIInF8wPm+4AoT6L0FA1QWAad3OMdTRx4tFYlU= golang.org/x/net v0.55.0 h1:bcvxaJn3e1U6InsFWt1JUq1aSjnRxLzT2rtD2KfkDF8= golang.org/x/net v0.55.0/go.mod h1:L5U2KuzuOe1lY7Z+aWVIKK6qEeJXnXV9yzGA+WCHJww= golang.org/x/sync v0.21.0 h1:HLII4xRRTtCRkxYp4HNFF0Js/Og6q2i++KXbg0gHCwM= golang.org/x/sync v0.21.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= +golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.45.0 h1:dO4czNzziLiiXplLQgBCEpCvXQ3dnkn0SdaZSYdQ+FY= golang.org/x/sys v0.45.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= golang.org/x/text v0.37.0 h1:Cqjiwd9eSg8e0QAkyCaQTNHFIIzWtidPahFWR83rTrc= golang.org/x/text v0.37.0/go.mod h1:a5sjxXGs9hsn/AJVwuElvCAo9v8QYLzvavO5z2PiM38= +golang.org/x/tools v0.45.0 h1:18qN3FAooORvApf5XjCXgsuayZOEtXf6JK18I3+ONa8= +golang.org/x/tools v0.45.0/go.mod h1:LuUGqqaXcXMEFEruIVJVm5mgDD8vww/z/SR1gQ4uE/0= gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4= gonum.org/v1/gonum v0.17.0/go.mod h1:El3tOrEuMpv2UdMrbNlKEh9vd86bmQ6vqIcDwxEOc1E= google.golang.org/genproto/googleapis/api v0.0.0-20260226221140-a57be14db171 h1:tu/dtnW1o3wfaxCOjSLn5IRX4YDcJrtlpzYkhHhGaC4= @@ -130,3 +147,31 @@ gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c/go.mod h1:JHkPIbrfpd72SG/EV gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= +modernc.org/cc/v4 v4.27.1 h1:9W30zRlYrefrDV2JE2O8VDtJ1yPGownxciz5rrbQZis= +modernc.org/cc/v4 v4.27.1/go.mod h1:uVtb5OGqUKpoLWhqwNQo/8LwvoiEBLvZXIQ/SmO6mL0= +modernc.org/ccgo/v4 v4.32.0 h1:hjG66bI/kqIPX1b2yT6fr/jt+QedtP2fqojG2VrFuVw= +modernc.org/ccgo/v4 v4.32.0/go.mod h1:6F08EBCx5uQc38kMGl+0Nm0oWczoo1c7cgpzEry7Uc0= +modernc.org/fileutil v1.4.0 h1:j6ZzNTftVS054gi281TyLjHPp6CPHr2KCxEXjEbD6SM= +modernc.org/fileutil v1.4.0/go.mod h1:EqdKFDxiByqxLk8ozOxObDSfcVOv/54xDs/DUHdvCUU= +modernc.org/gc/v2 v2.6.5 h1:nyqdV8q46KvTpZlsw66kWqwXRHdjIlJOhG6kxiV/9xI= +modernc.org/gc/v2 v2.6.5/go.mod h1:YgIahr1ypgfe7chRuJi2gD7DBQiKSLMPgBQe9oIiito= +modernc.org/gc/v3 v3.1.2 h1:ZtDCnhonXSZexk/AYsegNRV1lJGgaNZJuKjJSWKyEqo= +modernc.org/gc/v3 v3.1.2/go.mod h1:HFK/6AGESC7Ex+EZJhJ2Gni6cTaYpSMmU/cT9RmlfYY= +modernc.org/goabi0 v0.2.0 h1:HvEowk7LxcPd0eq6mVOAEMai46V+i7Jrj13t4AzuNks= +modernc.org/goabi0 v0.2.0/go.mod h1:CEFRnnJhKvWT1c1JTI3Avm+tgOWbkOu5oPA8eH8LnMI= +modernc.org/libc v1.70.0 h1:U58NawXqXbgpZ/dcdS9kMshu08aiA6b7gusEusqzNkw= +modernc.org/libc v1.70.0/go.mod h1:OVmxFGP1CI/Z4L3E0Q3Mf1PDE0BucwMkcXjjLntvHJo= +modernc.org/mathutil v1.7.1 h1:GCZVGXdaN8gTqB1Mf/usp1Y/hSqgI2vAGGP4jZMCxOU= +modernc.org/mathutil v1.7.1/go.mod h1:4p5IwJITfppl0G4sUEDtCr4DthTaT47/N3aT6MhfgJg= +modernc.org/memory v1.11.0 h1:o4QC8aMQzmcwCK3t3Ux/ZHmwFPzE6hf2Y5LbkRs+hbI= +modernc.org/memory v1.11.0/go.mod h1:/JP4VbVC+K5sU2wZi9bHoq2MAkCnrt2r98UGeSK7Mjw= +modernc.org/opt v0.1.4 h1:2kNGMRiUjrp4LcaPuLY2PzUfqM/w9N23quVwhKt5Qm8= +modernc.org/opt v0.1.4/go.mod h1:03fq9lsNfvkYSfxrfUhZCWPk1lm4cq4N+Bh//bEtgns= +modernc.org/sortutil v1.2.1 h1:+xyoGf15mM3NMlPDnFqrteY07klSFxLElE2PVuWIJ7w= +modernc.org/sortutil v1.2.1/go.mod h1:7ZI3a3REbai7gzCLcotuw9AC4VZVpYMjDzETGsSMqJE= +modernc.org/sqlite v1.48.1 h1:S85iToyU6cgeojybE2XJlSbcsvcWkQ6qqNXJHtW5hWA= +modernc.org/sqlite v1.48.1/go.mod h1:hWjRO6Tj/5Ik8ieqxQybiEOUXy0NJFNp2tpvVpKlvig= +modernc.org/strutil v1.2.1 h1:UneZBkQA+DX2Rp35KcM69cSsNES9ly8mQWD71HKlOA0= +modernc.org/strutil v1.2.1/go.mod h1:EHkiggD70koQxjVdSBM3JKM7k6L0FbGE5eymy9i3B9A= +modernc.org/token v1.1.0 h1:Xl7Ap9dKaEs5kLoOQeQmPWevfnk/DM5qcLcYlA8ys6Y= +modernc.org/token v1.1.0/go.mod h1:UGzOrNV1mAFSEB63lOFHIpNRUVMvYTc6yu1SMY/XTDM= diff --git a/internal/cmd/shim.go b/internal/cmd/shim.go index 654500429a..5361849ca6 100644 --- a/internal/cmd/shim.go +++ b/internal/cmd/shim.go @@ -4,6 +4,7 @@ import ( "github.com/sqlc-dev/sqlc/internal/compiler" "github.com/sqlc-dev/sqlc/internal/config" "github.com/sqlc-dev/sqlc/internal/config/convert" + "github.com/sqlc-dev/sqlc/internal/core" "github.com/sqlc-dev/sqlc/internal/info" "github.com/sqlc-dev/sqlc/internal/plugin" "github.com/sqlc-dev/sqlc/internal/sql/catalog" @@ -224,10 +225,55 @@ func pluginQueryParam(p compiler.Parameter) *plugin.Parameter { } func codeGenRequest(r *compiler.Result, settings config.CombinedSettings) *plugin.GenerateRequest { + var cat *plugin.Catalog + if r.CoreCatalog != nil { + cat = pluginCatalogFromCore(r.CoreCatalog) + } else { + cat = pluginCatalog(r.Catalog) + } return &plugin.GenerateRequest{ Settings: pluginSettings(r, settings), - Catalog: pluginCatalog(r.Catalog), + Catalog: cat, Queries: pluginQueries(r), SqlcVersion: info.Version, } } + +func pluginCatalogFromCore(cc *core.Catalog) *plugin.Catalog { + var schemas []*plugin.Schema + + namespaces, err := cc.Namespaces() + if err != nil { + return &plugin.Catalog{DefaultSchema: "public"} + } + for _, ns := range namespaces { + tables, err := cc.TablesInNamespace(ns.OID) + if err != nil { + continue + } + var ptables []*plugin.Table + for _, cl := range tables { + rel := &plugin.Identifier{Schema: ns.Name, Name: cl.Name} + cols, err := cc.ClassCodegenColumns(cl.OID) + if err != nil { + continue + } + var columns []*plugin.Column + for _, col := range cols { + columns = append(columns, &plugin.Column{ + Name: col.Name, + Type: &plugin.Identifier{Name: col.TypeName}, + NotNull: col.NotNull, + Table: rel, + }) + } + ptables = append(ptables, &plugin.Table{Rel: rel, Columns: columns}) + } + schemas = append(schemas, &plugin.Schema{Name: ns.Name, Tables: ptables}) + } + + return &plugin.Catalog{ + DefaultSchema: "public", + Schemas: schemas, + } +} diff --git a/internal/codegen/golang/clickhouse_type.go b/internal/codegen/golang/clickhouse_type.go new file mode 100644 index 0000000000..f1550b3534 --- /dev/null +++ b/internal/codegen/golang/clickhouse_type.go @@ -0,0 +1,59 @@ +package golang + +import ( + "strings" + + "github.com/sqlc-dev/sqlc/internal/codegen/golang/opts" + "github.com/sqlc-dev/sqlc/internal/codegen/sdk" + "github.com/sqlc-dev/sqlc/internal/plugin" +) + +func clickhouseType(req *plugin.GenerateRequest, options *opts.Options, col *plugin.Column) string { + dt := strings.ToLower(sdk.DataType(col.Type)) + notNull := col.NotNull + + switch dt { + case "uint8": + return nullable(notNull, "uint8") + case "uint16": + return nullable(notNull, "uint16") + case "uint32": + return nullable(notNull, "uint32") + case "uint64": + return nullable(notNull, "uint64") + case "int8": + return nullable(notNull, "int8") + case "int16": + return nullable(notNull, "int16") + case "int32": + return nullable(notNull, "int32") + case "int64": + return nullable(notNull, "int64") + case "uint128", "uint256", "int128", "int256": + return "*big.Int" + case "float32", "bfloat16": + return nullable(notNull, "float32") + case "float64": + return nullable(notNull, "float64") + case "bool": + return nullable(notNull, "bool") + case "string", "fixedstring": + return nullable(notNull, "string") + case "date", "date32", "datetime", "datetime64": + return nullable(notNull, "time.Time") + + case "decimal", "decimal32", "decimal64", "decimal128", "decimal256", + "uuid", "ipv4", "ipv6", "json", "enum8", "enum16": + return nullable(notNull, "string") + + default: + return "interface{}" + } +} + +func nullable(notNull bool, base string) string { + if notNull { + return base + } + return "*" + base +} diff --git a/internal/codegen/golang/go_type.go b/internal/codegen/golang/go_type.go index 21914736d0..b203dd525d 100644 --- a/internal/codegen/golang/go_type.go +++ b/internal/codegen/golang/go_type.go @@ -86,6 +86,8 @@ func goInnerType(req *plugin.GenerateRequest, options *opts.Options, col *plugin return postgresType(req, options, col) case "sqlite": return sqliteType(req, options, col) + case "clickhouse": + return clickhouseType(req, options, col) default: return "interface{}" } diff --git a/internal/compiler/compile.go b/internal/compiler/compile.go index b6bba42e16..534521c7a2 100644 --- a/internal/compiler/compile.go +++ b/internal/compiler/compile.go @@ -9,6 +9,7 @@ import ( "path/filepath" "strings" + "github.com/sqlc-dev/sqlc/internal/engine/clickhouse" "github.com/sqlc-dev/sqlc/internal/migrations" "github.com/sqlc-dev/sqlc/internal/multierr" "github.com/sqlc-dev/sqlc/internal/opts" @@ -55,6 +56,16 @@ func (c *Compiler) parseCatalog(schemas []string) error { continue } + if c.coreCatalog != nil { + for i := range stmts { + if err := clickhouse.Apply(c.coreCatalog, stmts[i].Raw); err != nil { + merr.Add(filename, contents, stmts[i].Pos(), err) + continue + } + } + continue + } + for i := range stmts { if err := c.catalog.Update(stmts[i], c); err != nil { merr.Add(filename, contents, stmts[i].Pos(), err) @@ -135,7 +146,8 @@ func (c *Compiler) parseQueries(o opts.Parser) (*Result, error) { } return &Result{ - Catalog: c.catalog, - Queries: q, + Catalog: c.catalog, + CoreCatalog: c.coreCatalog, + Queries: q, }, nil } diff --git a/internal/compiler/engine.go b/internal/compiler/engine.go index 64fdf3d5c7..8b50d30426 100644 --- a/internal/compiler/engine.go +++ b/internal/compiler/engine.go @@ -6,7 +6,9 @@ import ( "github.com/sqlc-dev/sqlc/internal/analyzer" "github.com/sqlc-dev/sqlc/internal/config" + "github.com/sqlc-dev/sqlc/internal/core" "github.com/sqlc-dev/sqlc/internal/dbmanager" + "github.com/sqlc-dev/sqlc/internal/engine/clickhouse" "github.com/sqlc-dev/sqlc/internal/engine/dolphin" "github.com/sqlc-dev/sqlc/internal/engine/postgresql" pganalyze "github.com/sqlc-dev/sqlc/internal/engine/postgresql/analyzer" @@ -27,6 +29,8 @@ type Compiler struct { client dbmanager.Client selector selector + coreCatalog *core.Catalog + schema []string // databaseOnlyMode indicates that the compiler should use database-only analysis @@ -111,6 +115,14 @@ func NewCompiler(conf config.SQL, combo config.CombinedSettings, parserOpts opts ) } } + case config.EngineClickHouse: + c.parser = clickhouse.NewParser() + c.selector = newDefaultSelector() + cat, err := core.New(clickhouse.Dialect()) + if err != nil { + return nil, fmt.Errorf("clickhouse: init catalog: %w", err) + } + c.coreCatalog = cat default: return nil, fmt.Errorf("unknown engine: %s", conf.Engine) } @@ -145,4 +157,7 @@ func (c *Compiler) Close(ctx context.Context) { if c.client != nil { c.client.Close(ctx) } + if c.coreCatalog != nil { + c.coreCatalog.Close() + } } diff --git a/internal/compiler/parse.go b/internal/compiler/parse.go index 2f9afb72c1..be510d71d4 100644 --- a/internal/compiler/parse.go +++ b/internal/compiler/parse.go @@ -19,6 +19,10 @@ import ( var debugDumpAST = sqlcdebug.New("dumpast") func (c *Compiler) parseQuery(stmt ast.Node, src string, o opts.Parser) (*Query, error) { + if c.coreCatalog != nil { + return c.parseQueryCore(stmt, src) + } + ctx := context.Background() if debugDumpAST.Value() == "1" { diff --git a/internal/compiler/parse_clickhouse.go b/internal/compiler/parse_clickhouse.go new file mode 100644 index 0000000000..9d79b21709 --- /dev/null +++ b/internal/compiler/parse_clickhouse.go @@ -0,0 +1,114 @@ +package compiler + +import ( + "errors" + "strings" + + "github.com/sqlc-dev/sqlc/internal/core" + coreanalyzer "github.com/sqlc-dev/sqlc/internal/core/analyzer" + "github.com/sqlc-dev/sqlc/internal/metadata" + "github.com/sqlc-dev/sqlc/internal/source" + "github.com/sqlc-dev/sqlc/internal/sql/ast" + "github.com/sqlc-dev/sqlc/internal/sql/validate" +) + +func (c *Compiler) parseQueryCore(stmt ast.Node, src string) (*Query, error) { + raw, ok := stmt.(*ast.RawStmt) + if !ok { + return nil, errors.New("node is not a statement") + } + rawSQL, err := source.Pluck(src, raw.StmtLocation, raw.StmtLen) + if err != nil { + return nil, err + } + if strings.TrimSpace(rawSQL) == "" { + return nil, errors.New("missing semicolon at end of file") + } + + name, cmd, err := metadata.ParseQueryNameAndType(rawSQL, metadata.CommentSyntax(c.parser.CommentSyntax())) + if err != nil { + return nil, err + } + if name == "" { + return nil, nil + } + if err := validate.Cmd(raw.Stmt, name, cmd); err != nil { + return nil, err + } + + md := metadata.Metadata{Name: name, Cmd: cmd} + cleanedComments, err := source.CleanedComments(rawSQL, c.parser.CommentSyntax()) + if err != nil { + return nil, err + } + md.Params, md.Flags, md.RuleSkiplist, err = metadata.ParseCommentFlags(cleanedComments) + if err != nil { + return nil, err + } + + var cols []*Column + var params []Parameter + if _, ok := raw.Stmt.(*ast.SelectStmt); ok { + res, err := coreanalyzer.Prepare(c.coreCatalog, raw) + if err != nil { + return nil, err + } + for _, col := range res.Columns { + cols = append(cols, coreColumn(col)) + } + for _, p := range res.Parameters { + params = append(params, Parameter{Number: p.Number, Column: coreParamColumn(p)}) + } + } + + trimmed, comments, err := source.StripComments(rawSQL) + if err != nil { + return nil, err + } + md.Comments = comments + + var insertTable *ast.TableName + if ins, ok := raw.Stmt.(*ast.InsertStmt); ok { + insertTable, _ = ParseTableName(ins.Relation) + } + + return &Query{ + RawStmt: raw, + Metadata: md, + Params: params, + Columns: cols, + SQL: trimmed, + InsertIntoTable: insertTable, + }, nil +} + +func coreColumn(c core.Column) *Column { + col := &Column{ + Name: c.Name, + DataType: c.DataType, + NotNull: c.NotNull, + } + if c.Source != nil && c.Source.Table != "" { + col.Table = &ast.TableName{Schema: c.Source.Schema, Name: c.Source.Table} + col.TableAlias = c.Source.TableAlias + col.OriginalName = c.Source.Column + } + if c.TypeLength > 0 { + l := c.TypeLength + col.Length = &l + } + return col +} + +func coreParamColumn(p core.Parameter) *Column { + col := &Column{ + Name: p.Name, + DataType: p.DataType, + NotNull: p.NotNull, + } + if p.Source != nil && p.Source.Table != "" { + col.Table = &ast.TableName{Schema: p.Source.Schema, Name: p.Source.Table} + col.OriginalName = p.Source.Column + } + return col +} diff --git a/internal/compiler/result.go b/internal/compiler/result.go index 3647da630f..69d78d8923 100644 --- a/internal/compiler/result.go +++ b/internal/compiler/result.go @@ -1,10 +1,13 @@ package compiler import ( + "github.com/sqlc-dev/sqlc/internal/core" "github.com/sqlc-dev/sqlc/internal/sql/catalog" ) type Result struct { Catalog *catalog.Catalog Queries []*Query + + CoreCatalog *core.Catalog } diff --git a/internal/config/config.go b/internal/config/config.go index ff7faedcaa..6f0a9bf0eb 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -54,6 +54,7 @@ const ( EngineMySQL Engine = "mysql" EnginePostgreSQL Engine = "postgresql" EngineSQLite Engine = "sqlite" + EngineClickHouse Engine = "clickhouse" ) type Config struct { diff --git a/internal/core/analysis.go b/internal/core/analysis.go new file mode 100644 index 0000000000..b437c46609 --- /dev/null +++ b/internal/core/analysis.go @@ -0,0 +1,48 @@ +package core + +type Command string + +const ( + CommandSelect Command = "SELECT" + CommandInsert Command = "INSERT" + CommandUpdate Command = "UPDATE" + CommandDelete Command = "DELETE" +) + +type PrepareResult struct { + Command Command `json:"command,omitempty"` + Columns []Column `json:"columns"` + Parameters []Parameter `json:"parameters"` +} + +type ColumnSource struct { + Schema string `json:"schema,omitempty"` + Table string `json:"table,omitempty"` + TableAlias string `json:"table_alias,omitempty"` + Column string `json:"column,omitempty"` +} + +type Column struct { + Name string `json:"name"` + DataType string `json:"data_type"` + TypeOID int64 `json:"type_oid,omitempty"` + NotNull bool `json:"not_null"` + SourceClassOID int64 `json:"source_class_oid,omitempty"` + SourceAttributeOID int64 `json:"source_attribute_oid,omitempty"` + Source *ColumnSource `json:"source,omitempty"` + DeclType string `json:"decl_type,omitempty"` + TypeLength int `json:"type_length,omitempty"` + TypeScale int `json:"type_scale,omitempty"` + IsPrimaryKey bool `json:"is_primary_key,omitempty"` + IsUnique bool `json:"is_unique,omitempty"` + IsAutoIncrement bool `json:"is_auto_increment,omitempty"` +} + +type Parameter struct { + Number int `json:"number"` + Name string `json:"name,omitempty"` + DataType string `json:"data_type,omitempty"` + TypeOID int64 `json:"type_oid,omitempty"` + NotNull bool `json:"not_null"` + Source *ColumnSource `json:"source,omitempty"` +} diff --git a/internal/core/analyzer/analyzer.go b/internal/core/analyzer/analyzer.go new file mode 100644 index 0000000000..d4fcc57bbc --- /dev/null +++ b/internal/core/analyzer/analyzer.go @@ -0,0 +1,136 @@ +package analyzer + +import ( + "fmt" + + "github.com/sqlc-dev/sqlc/internal/core" + "github.com/sqlc-dev/sqlc/internal/sql/ast" +) + +func Prepare(cat *core.Catalog, stmt ast.Node) (core.PrepareResult, error) { + if rs, ok := stmt.(*ast.RawStmt); ok { + stmt = rs.Stmt + } + a := &analyzer{ + cat: cat, + params: map[int]*core.Parameter{}, + } + switch s := stmt.(type) { + case *ast.SelectStmt: + if err := a.analyzeSelect(s); err != nil { + return core.PrepareResult{}, err + } + a.command = core.CommandSelect + default: + return core.PrepareResult{}, fmt.Errorf("analyzer: unsupported statement %T", stmt) + } + return a.result(), nil +} + +type analyzer struct { + cat *core.Catalog + scope *scope + columns []core.Column + params map[int]*core.Parameter + command core.Command +} + +func (a *analyzer) result() core.PrepareResult { + return core.PrepareResult{ + Command: a.command, + Columns: a.columns, + Parameters: orderedParams(a.params), + } +} + +func orderedParams(m map[int]*core.Parameter) []core.Parameter { + if len(m) == 0 { + return nil + } + maxN := 0 + for n := range m { + if n > maxN { + maxN = n + } + } + out := make([]core.Parameter, 0, len(m)) + for i := 1; i <= maxN; i++ { + if p, ok := m[i]; ok { + out = append(out, *p) + } + } + return out +} + +func (a *analyzer) analyzeSelect(s *ast.SelectStmt) error { + sc, err := a.buildScope(s.FromClause) + if err != nil { + return err + } + a.scope = sc + + for _, item := range listItems(s.FromClause) { + if err := a.typeJoinConditions(item); err != nil { + return fmt.Errorf("join: %w", err) + } + } + + if s.WhereClause != nil { + if _, err := a.typeExpr(s.WhereClause); err != nil { + return fmt.Errorf("where: %w", err) + } + } + if items := listItems(s.GroupClause); items != nil { + for _, g := range items { + if _, err := a.typeExpr(g); err != nil { + return fmt.Errorf("group by: %w", err) + } + } + } + if s.HavingClause != nil { + if _, err := a.typeExpr(s.HavingClause); err != nil { + return fmt.Errorf("having: %w", err) + } + } + + targets := listItems(s.TargetList) + if targets == nil { + return fmt.Errorf("select: empty target list") + } + for _, t := range targets { + rt, ok := t.(*ast.ResTarget) + if !ok { + continue + } + if err := a.projectTarget(rt); err != nil { + return err + } + } + return nil +} + +func listItems(l *ast.List) []ast.Node { + if l == nil { + return nil + } + return l.Items +} + +func (a *analyzer) typeJoinConditions(item ast.Node) error { + je, ok := item.(*ast.JoinExpr) + if !ok { + return nil + } + if err := a.typeJoinConditions(je.Larg); err != nil { + return err + } + if err := a.typeJoinConditions(je.Rarg); err != nil { + return err + } + if je.Quals != nil { + if _, err := a.typeExpr(je.Quals); err != nil { + return fmt.Errorf("ON: %w", err) + } + } + return nil +} diff --git a/internal/core/analyzer/analyzer_test.go b/internal/core/analyzer/analyzer_test.go new file mode 100644 index 0000000000..22204e9af9 --- /dev/null +++ b/internal/core/analyzer/analyzer_test.go @@ -0,0 +1,100 @@ +package analyzer_test + +import ( + "strings" + "testing" + + "github.com/sqlc-dev/sqlc/internal/core" + "github.com/sqlc-dev/sqlc/internal/core/analyzer" + "github.com/sqlc-dev/sqlc/internal/engine/postgresql" +) + +func seedUsers(t *testing.T) *core.Catalog { + t.Helper() + cat, err := core.New() + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { cat.Close() }) + + ns, err := cat.NamespaceOID("public") + if err != nil { + t.Fatal(err) + } + int4, err := cat.CreateType("int4", 4) + if err != nil { + t.Fatal(err) + } + text, err := cat.CreateType("text", -1) + if err != nil { + t.Fatal(err) + } + users, err := cat.CreateClass(ns, "users", "r") + if err != nil { + t.Fatal(err) + } + if err := cat.CreateAttribute(users, "id", int4, true, false, 1); err != nil { + t.Fatal(err) + } + if err := cat.CreateAttribute(users, "name", text, true, false, 2); err != nil { + t.Fatal(err) + } + return cat +} + +func prepare(t *testing.T, cat *core.Catalog, query string) core.PrepareResult { + t.Helper() + stmts, err := postgresql.NewParser().Parse(strings.NewReader(query)) + if err != nil { + t.Fatalf("parse: %v", err) + } + if len(stmts) != 1 { + t.Fatalf("expected 1 stmt, got %d", len(stmts)) + } + res, err := analyzer.Prepare(cat, stmts[0].Raw) + if err != nil { + t.Fatalf("analyze: %v", err) + } + return res +} + +func TestPrepareSelectColumns(t *testing.T) { + cat := seedUsers(t) + res := prepare(t, cat, "SELECT id, name FROM users") + + if len(res.Columns) != 2 { + t.Fatalf("got %d cols, want 2: %+v", len(res.Columns), res.Columns) + } + if res.Columns[0].Name != "id" || res.Columns[0].DataType != "int4" || !res.Columns[0].NotNull { + t.Errorf("col 0: %+v", res.Columns[0]) + } + if res.Columns[1].Name != "name" || res.Columns[1].DataType != "text" || !res.Columns[1].NotNull { + t.Errorf("col 1: %+v", res.Columns[1]) + } + for i, c := range res.Columns { + if c.SourceClassOID == 0 || c.SourceAttributeOID == 0 { + t.Errorf("col %d %s missing source binding: %+v", i, c.Name, c) + } + } +} + +func TestPrepareSelectStar(t *testing.T) { + cat := seedUsers(t) + res := prepare(t, cat, "SELECT * FROM users") + + if len(res.Columns) != 2 { + t.Fatalf("got %d cols, want 2: %+v", len(res.Columns), res.Columns) + } + if res.Columns[0].Name != "id" || res.Columns[1].Name != "name" { + t.Errorf("got %q, %q; want id, name", res.Columns[0].Name, res.Columns[1].Name) + } +} + +func TestPrepareAlias(t *testing.T) { + cat := seedUsers(t) + res := prepare(t, cat, "SELECT id AS user_id FROM users u") + + if len(res.Columns) != 1 || res.Columns[0].Name != "user_id" { + t.Fatalf("got %+v", res.Columns) + } +} diff --git a/internal/core/analyzer/expr.go b/internal/core/analyzer/expr.go new file mode 100644 index 0000000000..ff9ecd7e9d --- /dev/null +++ b/internal/core/analyzer/expr.go @@ -0,0 +1,325 @@ +package analyzer + +import ( + "fmt" + "strings" + + "github.com/sqlc-dev/sqlc/internal/core" + "github.com/sqlc-dev/sqlc/internal/sql/ast" +) + +type exprType struct { + typeOID int64 + nullable bool + sourceClassOID int64 + sourceAttributeOID int64 + sourceTableAlias string +} + +func (a *analyzer) typeExpr(n ast.Node) (exprType, error) { + switch e := n.(type) { + case nil: + return exprType{}, nil + + case *ast.TODO: + return exprType{}, nil + + case *ast.A_Const: + return a.typeConst(e) + + case *ast.ColumnRef: + return a.typeColumnRef(e) + + case *ast.ParamRef: + return a.typeParamRef(e) + + case *ast.A_Expr: + return a.typeAExpr(e) + + case *ast.BoolExpr: + return a.typeBoolExpr(e) + + case *ast.FuncCall: + return a.typeFuncCall(e) + + case *ast.TypeCast: + return a.typeTypeCast(e) + + case *ast.NullTest: + if _, err := a.typeExpr(e.Arg); err != nil { + return exprType{}, err + } + return a.boolType(false) + } + return exprType{}, fmt.Errorf("typeExpr: unsupported %T", n) +} + +func (a *analyzer) typeConst(c *ast.A_Const) (exprType, error) { + switch v := c.Val.(type) { + case *ast.Integer: + _ = v + oid, err := a.cat.TypeOID("int4") + if err != nil { + return exprType{}, err + } + return exprType{typeOID: oid}, nil + case *ast.Float: + oid, err := a.cat.TypeOID("numeric") + if err != nil { + return exprType{}, err + } + return exprType{typeOID: oid}, nil + case *ast.String: + oid, err := a.cat.TypeOID("text") + if err != nil { + return exprType{}, err + } + return exprType{typeOID: oid}, nil + case *ast.Boolean: + return a.boolType(false) + case nil: + return exprType{nullable: true}, nil + } + return exprType{}, fmt.Errorf("typeConst: unsupported %T", c.Val) +} + +func (a *analyzer) boolType(nullable bool) (exprType, error) { + oid, err := a.cat.TypeOID("bool") + if err != nil { + return exprType{}, err + } + return exprType{typeOID: oid, nullable: nullable}, nil +} + +func (a *analyzer) typeColumnRef(c *ast.ColumnRef) (exprType, error) { + parts := flattenFields(c.Fields) + if len(parts) == 0 { + return exprType{}, fmt.Errorf("column ref: empty") + } + relation := "" + column := parts[0] + if len(parts) >= 2 { + relation = parts[0] + column = parts[1] + } + rel, col, ok, err := a.scope.resolveColumn(relation, column) + if err != nil { + return exprType{}, err + } + if !ok { + if relation != "" { + return exprType{}, fmt.Errorf("unknown column %q.%q", relation, column) + } + return exprType{}, fmt.Errorf("unknown column %q", column) + } + return exprType{ + typeOID: col.typeOID, + nullable: !col.notNull, + sourceClassOID: rel.classOID, + sourceAttributeOID: col.attOID, + sourceTableAlias: rel.alias, + }, nil +} + +func flattenFields(fields *ast.List) []string { + if fields == nil { + return nil + } + out := make([]string, 0, len(fields.Items)) + for _, item := range fields.Items { + switch v := item.(type) { + case *ast.String: + out = append(out, v.Str) + case *ast.A_Star: + out = append(out, "*") + return out + } + } + return out +} + +func (a *analyzer) typeParamRef(p *ast.ParamRef) (exprType, error) { + cur, ok := a.params[p.Number] + if !ok { + cur = &core.Parameter{Number: p.Number} + a.params[p.Number] = cur + } + return exprType{typeOID: cur.TypeOID, nullable: !cur.NotNull}, nil +} + +func (a *analyzer) inferParam(number int, t exprType) { + cur, ok := a.params[number] + if !ok { + cur = &core.Parameter{Number: number} + a.params[number] = cur + } + if cur.TypeOID == 0 && t.typeOID != 0 { + cur.TypeOID = t.typeOID + if name, err := a.cat.TypeName(t.typeOID); err == nil { + cur.DataType = name + } + cur.NotNull = !t.nullable + } + if cur.Source == nil && t.sourceAttributeOID != 0 { + ad, err := a.cat.LookupAttribute(t.sourceAttributeOID) + if err == nil { + cur.Source = &core.ColumnSource{ + Schema: ad.Schema, + Table: ad.Table, + TableAlias: t.sourceTableAlias, + Column: ad.Column, + } + } + } +} + +func (a *analyzer) typeAExpr(e *ast.A_Expr) (exprType, error) { + if e.Kind != ast.A_Expr_Kind_OP { + return exprType{}, fmt.Errorf("a_expr: unsupported kind %v", e.Kind) + } + opName := opNameFromList(e.Name) + if opName == "" { + return exprType{}, fmt.Errorf("a_expr: unnamed operator") + } + + leftT, err := a.typeExpr(e.Lexpr) + if err != nil { + return exprType{}, err + } + rightT, err := a.typeExpr(e.Rexpr) + if err != nil { + return exprType{}, err + } + + if pr, ok := e.Lexpr.(*ast.ParamRef); ok && rightT.typeOID != 0 { + a.inferParam(pr.Number, rightT) + leftT = rightT + } + if pr, ok := e.Rexpr.(*ast.ParamRef); ok && leftT.typeOID != 0 { + a.inferParam(pr.Number, leftT) + rightT = leftT + } + + overload, err := a.resolveOperator(opName, leftT.typeOID, rightT.typeOID) + if err != nil { + return exprType{}, err + } + return exprType{ + typeOID: overload.ResultTypeOID, + nullable: leftT.nullable || rightT.nullable, + }, nil +} + +func opNameFromList(l *ast.List) string { + if l == nil { + return "" + } + parts := make([]string, 0, len(l.Items)) + for _, item := range l.Items { + if s, ok := item.(*ast.String); ok { + parts = append(parts, s.Str) + } + } + return strings.Join(parts, ".") +} + +func (a *analyzer) resolveOperator(name string, leftOID, rightOID int64) (core.OperatorOverload, error) { + candidates, err := a.cat.FindOperators(name, leftOID, rightOID) + if err != nil { + return core.OperatorOverload{}, err + } + if len(candidates) > 0 { + return candidates[0], nil + } + + all, err := a.cat.FindOperators(name, 0, 0) + if err != nil { + return core.OperatorOverload{}, err + } + for _, ov := range all { + if leftOID != 0 && ov.LeftTypeOID != 0 && leftOID != ov.LeftTypeOID { + ok, _ := a.cat.CastAllowed(leftOID, ov.LeftTypeOID, "i") + if !ok { + continue + } + } + if rightOID != 0 && ov.RightTypeOID != 0 && rightOID != ov.RightTypeOID { + ok, _ := a.cat.CastAllowed(rightOID, ov.RightTypeOID, "i") + if !ok { + continue + } + } + if (leftOID == 0) != (ov.LeftTypeOID == 0) { + continue + } + if (rightOID == 0) != (ov.RightTypeOID == 0) { + continue + } + return ov, nil + } + return core.OperatorOverload{}, fmt.Errorf("no operator %q for (%d, %d)", name, leftOID, rightOID) +} + +func (a *analyzer) typeBoolExpr(b *ast.BoolExpr) (exprType, error) { + for _, item := range listItems(b.Args) { + if _, err := a.typeExpr(item); err != nil { + return exprType{}, err + } + } + return a.boolType(false) +} + +func (a *analyzer) typeFuncCall(f *ast.FuncCall) (exprType, error) { + name := funcCallName(f) + if name == "" { + return exprType{}, fmt.Errorf("func call: missing name") + } + + if f.AggStar && (name == "count" || name == "count.*") { + oid, err := a.cat.TypeOID("int8") + if err != nil { + return exprType{}, err + } + return exprType{typeOID: oid, nullable: false}, nil + } + + for _, arg := range listItems(f.Args) { + if _, err := a.typeExpr(arg); err != nil { + return exprType{}, err + } + } + + overloads, err := a.cat.FindProcs(name, nil) + if err != nil { + return exprType{}, err + } + if len(overloads) == 0 { + return exprType{}, fmt.Errorf("unknown function %q", name) + } + p := overloads[0] + return exprType{typeOID: p.ReturnTypeOID, nullable: p.ReturnNullable}, nil +} + +func funcCallName(f *ast.FuncCall) string { + if f.Funcname != nil { + return opNameFromList(f.Funcname) + } + if f.Func != nil { + return strings.ToLower(f.Func.Name) + } + return "" +} + +func (a *analyzer) typeTypeCast(c *ast.TypeCast) (exprType, error) { + if c.TypeName == nil { + return exprType{}, fmt.Errorf("cast: missing target type") + } + if _, err := a.typeExpr(c.Arg); err != nil { + return exprType{}, err + } + oid, err := a.cat.TypeOID(strings.ToLower(c.TypeName.Name)) + if err != nil { + return exprType{}, fmt.Errorf("cast target %q: %w", c.TypeName.Name, err) + } + return exprType{typeOID: oid}, nil +} diff --git a/internal/core/analyzer/projection.go b/internal/core/analyzer/projection.go new file mode 100644 index 0000000000..3ed4034a76 --- /dev/null +++ b/internal/core/analyzer/projection.go @@ -0,0 +1,107 @@ +package analyzer + +import ( + "github.com/sqlc-dev/sqlc/internal/core" + "github.com/sqlc-dev/sqlc/internal/sql/ast" +) + +func (a *analyzer) projectTarget(rt *ast.ResTarget) error { + if cr, ok := rt.Val.(*ast.ColumnRef); ok { + if isStarRef(cr) { + a.emitStar(cr) + return nil + } + } + + t, err := a.typeExpr(rt.Val) + if err != nil { + return err + } + col := core.Column{ + Name: targetName(rt), + TypeOID: t.typeOID, + NotNull: !t.nullable, + SourceClassOID: t.sourceClassOID, + SourceAttributeOID: t.sourceAttributeOID, + } + if t.typeOID != 0 { + if name, err := a.cat.TypeName(t.typeOID); err == nil { + col.DataType = name + } + } + a.decorateSource(&col, t.sourceAttributeOID, t.sourceTableAlias) + a.columns = append(a.columns, col) + return nil +} + +func (a *analyzer) decorateSource(col *core.Column, attOID int64, tableAlias string) { + if attOID == 0 { + return + } + ad, err := a.cat.LookupAttribute(attOID) + if err != nil { + return + } + col.Source = &core.ColumnSource{ + Schema: ad.Schema, + Table: ad.Table, + TableAlias: tableAlias, + Column: ad.Column, + } + col.DeclType = ad.DeclType + col.TypeLength = ad.TypeLength + col.TypeScale = ad.TypeScale + col.IsPrimaryKey = ad.IsPrimaryKey + col.IsUnique = ad.IsUnique + col.IsAutoIncrement = ad.AutoIncrement +} + +func targetName(rt *ast.ResTarget) string { + if rt.Name != nil && *rt.Name != "" { + return *rt.Name + } + if cr, ok := rt.Val.(*ast.ColumnRef); ok { + parts := flattenFields(cr.Fields) + if len(parts) > 0 { + return parts[len(parts)-1] + } + } + if fc, ok := rt.Val.(*ast.FuncCall); ok { + if name := funcCallName(fc); name != "" { + return name + } + } + return "?column?" +} + +func isStarRef(c *ast.ColumnRef) bool { + parts := flattenFields(c.Fields) + return len(parts) > 0 && parts[len(parts)-1] == "*" +} + +func (a *analyzer) emitStar(cr *ast.ColumnRef) { + parts := flattenFields(cr.Fields) + relName := "" + if len(parts) > 1 { + relName = parts[0] + } + for _, rel := range a.scope.rels { + if relName != "" && rel.alias != relName { + continue + } + for _, c := range rel.cols { + col := core.Column{ + Name: c.name, + TypeOID: c.typeOID, + NotNull: c.notNull, + SourceClassOID: rel.classOID, + SourceAttributeOID: c.attOID, + } + if name, err := a.cat.TypeName(c.typeOID); err == nil { + col.DataType = name + } + a.decorateSource(&col, c.attOID, rel.alias) + a.columns = append(a.columns, col) + } + } +} diff --git a/internal/core/analyzer/scope.go b/internal/core/analyzer/scope.go new file mode 100644 index 0000000000..772d6e2dc5 --- /dev/null +++ b/internal/core/analyzer/scope.go @@ -0,0 +1,136 @@ +package analyzer + +import ( + "fmt" + + "github.com/sqlc-dev/sqlc/internal/core" + "github.com/sqlc-dev/sqlc/internal/sql/ast" +) + +type scope struct { + rels []scopeRel +} + +type scopeRel struct { + alias string + classOID int64 + cols []scopeCol +} + +type scopeCol struct { + name string + attOID int64 + typeOID int64 + notNull bool +} + +func (a *analyzer) buildScope(from *ast.List) (*scope, error) { + sc := &scope{} + for _, item := range listItems(from) { + if err := a.appendFromItem(sc, item); err != nil { + return nil, err + } + } + return sc, nil +} + +func (a *analyzer) appendFromItem(sc *scope, item ast.Node) error { + switch v := item.(type) { + case *ast.RangeVar: + rel, err := a.bindRangeVar(v) + if err != nil { + return err + } + sc.rels = append(sc.rels, rel) + return nil + case *ast.JoinExpr: + if err := a.appendFromItem(sc, v.Larg); err != nil { + return err + } + return a.appendFromItem(sc, v.Rarg) + default: + return fmt.Errorf("scope: unsupported FROM item %T", item) + } +} + +func (a *analyzer) bindRangeVar(rv *ast.RangeVar) (scopeRel, error) { + if rv.Relname == nil { + return scopeRel{}, fmt.Errorf("range var: missing relation name") + } + relName := *rv.Relname + schema := "" + if rv.Schemaname != nil { + schema = *rv.Schemaname + } + if schema == "" { + schema = "public" + } + nsOID, err := a.cat.NamespaceOID(schema) + if err != nil { + return scopeRel{}, fmt.Errorf("schema %q: %w", schema, err) + } + classOID, err := a.cat.ClassOID(nsOID, relName) + if err != nil { + return scopeRel{}, fmt.Errorf("relation %q.%q: %w", schema, relName, err) + } + rel := scopeRel{ + alias: relName, + classOID: classOID, + } + if rv.Alias != nil && rv.Alias.Aliasname != nil && *rv.Alias.Aliasname != "" { + rel.alias = *rv.Alias.Aliasname + } + + cols, err := a.classColumns(classOID) + if err != nil { + return scopeRel{}, err + } + rel.cols = cols + return rel, nil +} + +func (a *analyzer) classColumns(classOID int64) ([]scopeCol, error) { + cols, err := a.cat.ClassColumns(classOID) + if err != nil { + return nil, err + } + out := make([]scopeCol, 0, len(cols)) + for _, c := range cols { + out = append(out, scopeCol{ + name: c.Name, + attOID: c.AttOID, + typeOID: c.TypeOID, + notNull: c.NotNull, + }) + } + return out, nil +} + +func (s *scope) resolveColumn(relation, column string) (rel scopeRel, col scopeCol, ok bool, err error) { + var matches []struct { + rel scopeRel + col scopeCol + } + for _, r := range s.rels { + if relation != "" && r.alias != relation { + continue + } + for _, c := range r.cols { + if c.name == column { + matches = append(matches, struct { + rel scopeRel + col scopeCol + }{r, c}) + } + } + } + if len(matches) == 0 { + return scopeRel{}, scopeCol{}, false, nil + } + if len(matches) > 1 { + return scopeRel{}, scopeCol{}, false, fmt.Errorf("ambiguous column reference %q", column) + } + return matches[0].rel, matches[0].col, true, nil +} + +var _ = core.Column{} diff --git a/internal/core/attribute.go b/internal/core/attribute.go new file mode 100644 index 0000000000..599479c26c --- /dev/null +++ b/internal/core/attribute.go @@ -0,0 +1,228 @@ +package core + +import ( + "context" + "fmt" + + "github.com/sqlc-dev/sqlc/internal/core/catalogdb" +) + +type AttributeSpec struct { + ClassOID int64 + Name string + TypeOID int64 + Num int + NotNull bool + HasDefault bool + DeclType string + TypeLength int + TypeScale int + AutoIncrement bool + IsPrimaryKey bool + IsUnique bool +} + +func (c *Catalog) CreateAttributeSpec(s AttributeSpec) error { + err := c.q.CreateAttribute(context.Background(), catalogdb.CreateAttributeParams{ + ClassOid: s.ClassOID, + Name: s.Name, + TypeOid: s.TypeOID, + NotNull: boolToInt64(s.NotNull), + HasDefault: boolToInt64(s.HasDefault), + Num: int64(s.Num), + DeclType: s.DeclType, + TypeLength: int64(s.TypeLength), + TypeScale: int64(s.TypeScale), + AutoIncrement: boolToInt64(s.AutoIncrement), + IsPrimaryKey: boolToInt64(s.IsPrimaryKey), + IsUnique: boolToInt64(s.IsUnique), + }) + if err != nil { + return fmt.Errorf("create attribute %q on class %d: %w", s.Name, s.ClassOID, err) + } + return nil +} + +func (c *Catalog) CreateAttribute(classOID int64, name string, typeOID int64, notNull bool, hasDefault bool, num int) error { + return c.CreateAttributeSpec(AttributeSpec{ + ClassOID: classOID, + Name: name, + TypeOID: typeOID, + Num: num, + NotNull: notNull, + HasDefault: hasDefault, + }) +} + +func (c *Catalog) SetAttributePrimaryKey(classOID int64, columns []string) error { + ctx := context.Background() + for _, name := range columns { + err := c.q.SetAttributePrimaryKey(ctx, catalogdb.SetAttributePrimaryKeyParams{ + ClassOid: classOID, + Name: name, + }) + if err != nil { + return fmt.Errorf("mark pk %s on class %d: %w", name, classOID, err) + } + } + return nil +} + +func (c *Catalog) SetAttributeUnique(classOID int64, columns []string) error { + ctx := context.Background() + for _, name := range columns { + err := c.q.SetAttributeUnique(ctx, catalogdb.SetAttributeUniqueParams{ + ClassOid: classOID, + Name: name, + }) + if err != nil { + return fmt.Errorf("mark unique %s on class %d: %w", name, classOID, err) + } + } + return nil +} + +type ColumnInfo struct { + Name string + TypeName string + TypeOID int64 + NotNull bool + DeclType string + TypeLength int + TypeScale int + AutoIncrement bool + IsPrimaryKey bool + IsUnique bool + AttributeOID int64 + ClassOID int64 +} + +func (c *Catalog) ResolveColumn(table, column string) (*ColumnInfo, error) { + r, err := c.q.ResolveColumn(context.Background(), catalogdb.ResolveColumnParams{ + TableName: table, + ColumnName: column, + }) + if err != nil { + return nil, fmt.Errorf("resolve %s.%s: %w", table, column, err) + } + info := ColumnInfo{ + AttributeOID: r.Oid, + ClassOID: r.ClassOid, + Name: r.Name, + TypeName: r.TypeName, + TypeOID: r.TypeOid, + NotNull: r.NotNull != 0, + DeclType: r.DeclType, + TypeLength: int(r.TypeLength), + TypeScale: int(r.TypeScale), + AutoIncrement: r.AutoIncrement != 0, + IsPrimaryKey: r.IsPrimaryKey != 0, + IsUnique: r.IsUnique != 0, + } + return &info, nil +} + +func (c *Catalog) TableColumns(table string) ([]ColumnInfo, error) { + rows, err := c.q.TableColumns(context.Background(), table) + if err != nil { + return nil, fmt.Errorf("table columns %q: %w", table, err) + } + cols := make([]ColumnInfo, 0, len(rows)) + for _, r := range rows { + cols = append(cols, ColumnInfo{ + AttributeOID: r.Oid, + ClassOID: r.ClassOid, + Name: r.Name, + TypeName: r.TypeName, + TypeOID: r.TypeOid, + NotNull: r.NotNull != 0, + DeclType: r.DeclType, + TypeLength: int(r.TypeLength), + TypeScale: int(r.TypeScale), + AutoIncrement: r.AutoIncrement != 0, + IsPrimaryKey: r.IsPrimaryKey != 0, + IsUnique: r.IsUnique != 0, + }) + } + return cols, nil +} + +type ClassColumn struct { + AttOID int64 + Name string + TypeOID int64 + NotNull bool +} + +func (c *Catalog) ClassColumns(classOID int64) ([]ClassColumn, error) { + rows, err := c.q.ClassAttributes(context.Background(), classOID) + if err != nil { + return nil, fmt.Errorf("class columns %d: %w", classOID, err) + } + out := make([]ClassColumn, 0, len(rows)) + for _, r := range rows { + out = append(out, ClassColumn{ + AttOID: r.Oid, + Name: r.Name, + TypeOID: r.TypeOid, + NotNull: r.NotNull != 0, + }) + } + return out, nil +} + +type CodegenColumn struct { + Name string + TypeName string + NotNull bool +} + +func (c *Catalog) ClassCodegenColumns(classOID int64) ([]CodegenColumn, error) { + rows, err := c.q.ListClassColumns(context.Background(), classOID) + if err != nil { + return nil, fmt.Errorf("class codegen columns %d: %w", classOID, err) + } + out := make([]CodegenColumn, 0, len(rows)) + for _, r := range rows { + out = append(out, CodegenColumn{ + Name: r.ColumnName, + TypeName: r.TypeName, + NotNull: r.NotNull != 0, + }) + } + return out, nil +} + +type AttributeDetails struct { + Schema string + Table string + Column string + Num int + DeclType string + TypeLength int + TypeScale int + AutoIncrement bool + IsPrimaryKey bool + IsUnique bool + NotNull bool +} + +func (c *Catalog) LookupAttribute(attOID int64) (AttributeDetails, error) { + r, err := c.q.LookupAttribute(context.Background(), attOID) + if err != nil { + return AttributeDetails{}, fmt.Errorf("lookup attribute %d: %w", attOID, err) + } + return AttributeDetails{ + Schema: r.SchemaName, + Table: r.TableName, + Column: r.ColumnName, + Num: int(r.Num), + DeclType: r.DeclType, + TypeLength: int(r.TypeLength), + TypeScale: int(r.TypeScale), + AutoIncrement: r.AutoIncrement != 0, + IsPrimaryKey: r.IsPrimaryKey != 0, + IsUnique: r.IsUnique != 0, + NotNull: r.NotNull != 0, + }, nil +} diff --git a/internal/core/cast.go b/internal/core/cast.go new file mode 100644 index 0000000000..408663bf8e --- /dev/null +++ b/internal/core/cast.go @@ -0,0 +1,69 @@ +package core + +import ( + "context" + "fmt" + + "github.com/sqlc-dev/sqlc/internal/core/catalogdb" +) + +type CastSpec struct { + SourceTypeOID int64 + TargetTypeOID int64 + ProcOID int64 + Context string + DialectOID int64 +} + +func (c *Catalog) CreateCast(cs CastSpec) error { + if cs.Context == "" { + cs.Context = "e" + } + err := c.q.CreateCast(context.Background(), catalogdb.CreateCastParams{ + SourceTypeOid: cs.SourceTypeOID, + TargetTypeOid: cs.TargetTypeOID, + ProcOid: nullInt64(cs.ProcOID), + Context: cs.Context, + DialectOid: nullInt64(cs.DialectOID), + }) + if err != nil { + return fmt.Errorf("create cast %d->%d: %w", cs.SourceTypeOID, cs.TargetTypeOID, err) + } + return nil +} + +func (c *Catalog) FindCast(src, tgt int64) (CastSpec, bool, error) { + row, err := c.q.FindCast(context.Background(), catalogdb.FindCastParams{ + SourceTypeOid: src, + TargetTypeOid: tgt, + }) + if err != nil { + return CastSpec{}, false, nil + } + return CastSpec{ + SourceTypeOID: row.SourceTypeOid, + TargetTypeOID: row.TargetTypeOid, + ProcOID: orZero(row.ProcOid), + Context: row.Context, + DialectOID: orZero(row.DialectOid), + }, true, nil +} + +func (c *Catalog) CastAllowed(src, tgt int64, ctx string) (bool, error) { + if src == tgt { + return true, nil + } + cs, ok, err := c.FindCast(src, tgt) + if err != nil || !ok { + return false, err + } + switch ctx { + case "i": + return cs.Context == "i", nil + case "a": + return cs.Context == "i" || cs.Context == "a", nil + case "e": + return true, nil + } + return false, fmt.Errorf("unknown cast context %q", ctx) +} diff --git a/internal/core/cast_test.go b/internal/core/cast_test.go new file mode 100644 index 0000000000..8e1f7171ce --- /dev/null +++ b/internal/core/cast_test.go @@ -0,0 +1,51 @@ +package core + +import "testing" + +func TestCastAllowed(t *testing.T) { + cat, err := New() + if err != nil { + t.Fatal(err) + } + defer cat.Close() + + intOID, _ := cat.CreateType("integer", 4) + bigintOID, _ := cat.CreateType("bigint", 8) + textOID, _ := cat.CreateType("text", 0) + + if err := cat.CreateCast(CastSpec{ + SourceTypeOID: intOID, TargetTypeOID: bigintOID, Context: "i", + }); err != nil { + t.Fatal(err) + } + if err := cat.CreateCast(CastSpec{ + SourceTypeOID: intOID, TargetTypeOID: textOID, Context: "e", + }); err != nil { + t.Fatal(err) + } + + cases := []struct { + src, tgt int64 + ctx string + want bool + }{ + {intOID, intOID, "i", true}, + {intOID, bigintOID, "i", true}, + {intOID, bigintOID, "a", true}, + {intOID, bigintOID, "e", true}, + {intOID, textOID, "i", false}, + {intOID, textOID, "a", false}, + {intOID, textOID, "e", true}, + {textOID, intOID, "e", false}, + } + for _, c := range cases { + got, err := cat.CastAllowed(c.src, c.tgt, c.ctx) + if err != nil { + t.Errorf("CastAllowed(%d,%d,%q): %v", c.src, c.tgt, c.ctx, err) + continue + } + if got != c.want { + t.Errorf("CastAllowed(%d,%d,%q): got %v, want %v", c.src, c.tgt, c.ctx, got, c.want) + } + } +} diff --git a/internal/core/catalog.go b/internal/core/catalog.go new file mode 100644 index 0000000000..e87af5b7d6 --- /dev/null +++ b/internal/core/catalog.go @@ -0,0 +1,60 @@ +package core + +import ( + "database/sql" + "fmt" + + "github.com/sqlc-dev/sqlc/internal/core/catalogdb" + "github.com/sqlc-dev/sqlc/internal/core/catalogdef" + + _ "modernc.org/sqlite" +) + +//go:generate go run github.com/sqlc-dev/sqlc/cmd/sqlc generate + +type Catalog struct { + db *sql.DB + q *catalogdb.Queries +} + +type Option func(*Catalog) error + +func WithSeed(fn func(*Catalog) error) Option { + return func(c *Catalog) error { return fn(c) } +} + +func New(opts ...Option) (*Catalog, error) { + db, err := sql.Open("sqlite", ":memory:") + if err != nil { + return nil, fmt.Errorf("core: open catalog: %w", err) + } + if _, err := db.Exec(catalogdef.Schema); err != nil { + db.Close() + return nil, fmt.Errorf("core: init schema: %w", err) + } + c := &Catalog{db: db, q: catalogdb.New(db)} + if err := c.bootstrap(); err != nil { + db.Close() + return nil, fmt.Errorf("core: bootstrap: %w", err) + } + for i, opt := range opts { + if err := opt(c); err != nil { + db.Close() + return nil, fmt.Errorf("core: option %d: %w", i, err) + } + } + return c, nil +} + +func (c *Catalog) Close() error { + return c.db.Close() +} + +func (c *Catalog) DB() *sql.DB { + return c.db +} + +func (c *Catalog) bootstrap() error { + _, err := c.CreateNamespace("public") + return err +} diff --git a/internal/core/catalogdb/db.go b/internal/core/catalogdb/db.go new file mode 100644 index 0000000000..d16ad694e9 --- /dev/null +++ b/internal/core/catalogdb/db.go @@ -0,0 +1,31 @@ +// Code generated by sqlc. DO NOT EDIT. +// versions: +// sqlc v1.31.1 + +package catalogdb + +import ( + "context" + "database/sql" +) + +type DBTX interface { + ExecContext(context.Context, string, ...interface{}) (sql.Result, error) + PrepareContext(context.Context, string) (*sql.Stmt, error) + QueryContext(context.Context, string, ...interface{}) (*sql.Rows, error) + QueryRowContext(context.Context, string, ...interface{}) *sql.Row +} + +func New(db DBTX) *Queries { + return &Queries{db: db} +} + +type Queries struct { + db DBTX +} + +func (q *Queries) WithTx(tx *sql.Tx) *Queries { + return &Queries{ + db: tx, + } +} diff --git a/internal/core/catalogdb/models.go b/internal/core/catalogdb/models.go new file mode 100644 index 0000000000..374a2e70ee --- /dev/null +++ b/internal/core/catalogdb/models.go @@ -0,0 +1,111 @@ +// Code generated by sqlc. DO NOT EDIT. +// versions: +// sqlc v1.31.1 + +package catalogdb + +import ( + "database/sql" +) + +type SqlAttribute struct { + Oid int64 + ClassOid int64 + Name string + TypeOid int64 + NotNull int64 + HasDefault int64 + Num int64 + DeclType string + TypeLength int64 + TypeScale int64 + AutoIncrement int64 + IsPrimaryKey int64 + IsUnique int64 +} + +type SqlCast struct { + SourceTypeOid int64 + TargetTypeOid int64 + ProcOid sql.NullInt64 + Context string + DialectOid sql.NullInt64 +} + +type SqlClass struct { + Oid int64 + NamespaceOid int64 + Name string + Kind string +} + +type SqlConstraint struct { + Oid int64 + ClassOid int64 + Name string + Kind string + Columns string +} + +type SqlDialect struct { + Oid int64 + Name string +} + +type SqlDialectFlag struct { + DialectOid int64 + Key string + Value string +} + +type SqlNamespace struct { + Oid int64 + Name string +} + +type SqlOperator struct { + Oid int64 + NamespaceOid sql.NullInt64 + DialectOid sql.NullInt64 + Name string + LeftTypeOid sql.NullInt64 + RightTypeOid sql.NullInt64 + ResultTypeOid int64 + ProcOid sql.NullInt64 + CommutatorOid sql.NullInt64 + NegatorOid sql.NullInt64 +} + +type SqlProc struct { + Oid int64 + NamespaceOid sql.NullInt64 + DialectOid sql.NullInt64 + Name string + Kind string + ReturnTypeOid int64 + ReturnSet int64 + ReturnNullable int64 + Strict int64 + VariadicKind string +} + +type SqlProcArg struct { + ProcOid int64 + Ord int64 + Name string + TypeOid int64 + Mode string + HasDefault int64 +} + +type SqlType struct { + Oid int64 + NamespaceOid int64 + DialectOid sql.NullInt64 + Name string + Size int64 + Typtype string + Category sql.NullString + Preferred int64 + ElementOid sql.NullInt64 +} diff --git a/internal/core/catalogdb/query.sql.go b/internal/core/catalogdb/query.sql.go new file mode 100644 index 0000000000..3ec3eacc73 --- /dev/null +++ b/internal/core/catalogdb/query.sql.go @@ -0,0 +1,985 @@ +// Code generated by sqlc. DO NOT EDIT. +// versions: +// sqlc v1.31.1 +// source: query.sql + +package catalogdb + +import ( + "context" + "database/sql" + "strings" +) + +const classAttributes = `-- name: ClassAttributes :many +SELECT oid, name, type_oid, not_null +FROM sql_attribute +WHERE class_oid = ? +ORDER BY num +` + +type ClassAttributesRow struct { + Oid int64 + Name string + TypeOid int64 + NotNull int64 +} + +func (q *Queries) ClassAttributes(ctx context.Context, classOid int64) ([]ClassAttributesRow, error) { + rows, err := q.db.QueryContext(ctx, classAttributes, classOid) + if err != nil { + return nil, err + } + defer rows.Close() + var items []ClassAttributesRow + for rows.Next() { + var i ClassAttributesRow + if err := rows.Scan( + &i.Oid, + &i.Name, + &i.TypeOid, + &i.NotNull, + ); err != nil { + return nil, err + } + items = append(items, i) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const classOID = `-- name: ClassOID :one +SELECT oid FROM sql_class WHERE namespace_oid = ? AND name = ? +` + +type ClassOIDParams struct { + NamespaceOid int64 + Name string +} + +func (q *Queries) ClassOID(ctx context.Context, arg ClassOIDParams) (int64, error) { + row := q.db.QueryRowContext(ctx, classOID, arg.NamespaceOid, arg.Name) + var oid int64 + err := row.Scan(&oid) + return oid, err +} + +const classOIDByName = `-- name: ClassOIDByName :one +SELECT oid FROM sql_class WHERE name = ? LIMIT 1 +` + +func (q *Queries) ClassOIDByName(ctx context.Context, name string) (int64, error) { + row := q.db.QueryRowContext(ctx, classOIDByName, name) + var oid int64 + err := row.Scan(&oid) + return oid, err +} + +const createAttribute = `-- name: CreateAttribute :exec + +INSERT INTO sql_attribute ( + class_oid, name, type_oid, not_null, has_default, num, + decl_type, type_length, type_scale, + auto_increment, is_primary_key, is_unique +) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) +` + +type CreateAttributeParams struct { + ClassOid int64 + Name string + TypeOid int64 + NotNull int64 + HasDefault int64 + Num int64 + DeclType string + TypeLength int64 + TypeScale int64 + AutoIncrement int64 + IsPrimaryKey int64 + IsUnique int64 +} + +// ============================= sql_attribute =========================== +func (q *Queries) CreateAttribute(ctx context.Context, arg CreateAttributeParams) error { + _, err := q.db.ExecContext(ctx, createAttribute, + arg.ClassOid, + arg.Name, + arg.TypeOid, + arg.NotNull, + arg.HasDefault, + arg.Num, + arg.DeclType, + arg.TypeLength, + arg.TypeScale, + arg.AutoIncrement, + arg.IsPrimaryKey, + arg.IsUnique, + ) + return err +} + +const createCast = `-- name: CreateCast :exec + +INSERT INTO sql_cast (source_type_oid, target_type_oid, proc_oid, context, dialect_oid) +VALUES (?, ?, ?, ?, ?) +` + +type CreateCastParams struct { + SourceTypeOid int64 + TargetTypeOid int64 + ProcOid sql.NullInt64 + Context string + DialectOid sql.NullInt64 +} + +// =============================== sql_cast ============================== +func (q *Queries) CreateCast(ctx context.Context, arg CreateCastParams) error { + _, err := q.db.ExecContext(ctx, createCast, + arg.SourceTypeOid, + arg.TargetTypeOid, + arg.ProcOid, + arg.Context, + arg.DialectOid, + ) + return err +} + +const createClass = `-- name: CreateClass :execlastid + +INSERT INTO sql_class (namespace_oid, name, kind) VALUES (?, ?, ?) +` + +type CreateClassParams struct { + NamespaceOid int64 + Name string + Kind string +} + +// =============================== sql_class ============================= +func (q *Queries) CreateClass(ctx context.Context, arg CreateClassParams) (int64, error) { + result, err := q.db.ExecContext(ctx, createClass, arg.NamespaceOid, arg.Name, arg.Kind) + if err != nil { + return 0, err + } + return result.LastInsertId() +} + +const createConstraint = `-- name: CreateConstraint :exec + +INSERT INTO sql_constraint (class_oid, name, kind, columns) VALUES (?, ?, ?, ?) +` + +type CreateConstraintParams struct { + ClassOid int64 + Name string + Kind string + Columns string +} + +// ============================ sql_constraint =========================== +func (q *Queries) CreateConstraint(ctx context.Context, arg CreateConstraintParams) error { + _, err := q.db.ExecContext(ctx, createConstraint, + arg.ClassOid, + arg.Name, + arg.Kind, + arg.Columns, + ) + return err +} + +const createDialect = `-- name: CreateDialect :execlastid + +INSERT INTO sql_dialect (name) VALUES (?) +` + +// ============================== sql_dialect ============================ +func (q *Queries) CreateDialect(ctx context.Context, name string) (int64, error) { + result, err := q.db.ExecContext(ctx, createDialect, name) + if err != nil { + return 0, err + } + return result.LastInsertId() +} + +const createNamespace = `-- name: CreateNamespace :execlastid + + +INSERT INTO sql_namespace (name) VALUES (?) +` + +// Queries against sqlc's own sql_* catalog tables, compiled by sqlc's +// SQLite engine. Regenerate with `go generate ./internal/core/...`. +// ============================ sql_namespace ============================ +func (q *Queries) CreateNamespace(ctx context.Context, name string) (int64, error) { + result, err := q.db.ExecContext(ctx, createNamespace, name) + if err != nil { + return 0, err + } + return result.LastInsertId() +} + +const createOperator = `-- name: CreateOperator :execlastid + +INSERT INTO sql_operator + (namespace_oid, dialect_oid, name, + left_type_oid, right_type_oid, result_type_oid, proc_oid) +VALUES (?, ?, ?, ?, ?, ?, ?) +` + +type CreateOperatorParams struct { + NamespaceOid sql.NullInt64 + DialectOid sql.NullInt64 + Name string + LeftTypeOid sql.NullInt64 + RightTypeOid sql.NullInt64 + ResultTypeOid int64 + ProcOid sql.NullInt64 +} + +// ============================= sql_operator ============================ +func (q *Queries) CreateOperator(ctx context.Context, arg CreateOperatorParams) (int64, error) { + result, err := q.db.ExecContext(ctx, createOperator, + arg.NamespaceOid, + arg.DialectOid, + arg.Name, + arg.LeftTypeOid, + arg.RightTypeOid, + arg.ResultTypeOid, + arg.ProcOid, + ) + if err != nil { + return 0, err + } + return result.LastInsertId() +} + +const createProc = `-- name: CreateProc :execlastid + +INSERT INTO sql_proc + (namespace_oid, dialect_oid, name, kind, + return_type_oid, return_set, return_nullable, strict, variadic_kind) +VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) +` + +type CreateProcParams struct { + NamespaceOid sql.NullInt64 + DialectOid sql.NullInt64 + Name string + Kind string + ReturnTypeOid int64 + ReturnSet int64 + ReturnNullable int64 + Strict int64 + VariadicKind string +} + +// =============================== sql_proc ============================== +func (q *Queries) CreateProc(ctx context.Context, arg CreateProcParams) (int64, error) { + result, err := q.db.ExecContext(ctx, createProc, + arg.NamespaceOid, + arg.DialectOid, + arg.Name, + arg.Kind, + arg.ReturnTypeOid, + arg.ReturnSet, + arg.ReturnNullable, + arg.Strict, + arg.VariadicKind, + ) + if err != nil { + return 0, err + } + return result.LastInsertId() +} + +const createProcArg = `-- name: CreateProcArg :exec +INSERT INTO sql_proc_arg (proc_oid, ord, name, type_oid, mode, has_default) +VALUES (?, ?, ?, ?, ?, ?) +` + +type CreateProcArgParams struct { + ProcOid int64 + Ord int64 + Name string + TypeOid int64 + Mode string + HasDefault int64 +} + +func (q *Queries) CreateProcArg(ctx context.Context, arg CreateProcArgParams) error { + _, err := q.db.ExecContext(ctx, createProcArg, + arg.ProcOid, + arg.Ord, + arg.Name, + arg.TypeOid, + arg.Mode, + arg.HasDefault, + ) + return err +} + +const createType = `-- name: CreateType :execlastid + +INSERT INTO sql_type + (name, size, typtype, category, preferred, namespace_oid, dialect_oid, element_oid) +VALUES (?, ?, ?, ?, ?, ?, ?, ?) +` + +type CreateTypeParams struct { + Name string + Size int64 + Typtype string + Category sql.NullString + Preferred int64 + NamespaceOid int64 + DialectOid sql.NullInt64 + ElementOid sql.NullInt64 +} + +// =============================== sql_type ============================== +func (q *Queries) CreateType(ctx context.Context, arg CreateTypeParams) (int64, error) { + result, err := q.db.ExecContext(ctx, createType, + arg.Name, + arg.Size, + arg.Typtype, + arg.Category, + arg.Preferred, + arg.NamespaceOid, + arg.DialectOid, + arg.ElementOid, + ) + if err != nil { + return 0, err + } + return result.LastInsertId() +} + +const deleteAttributesByClass = `-- name: DeleteAttributesByClass :exec +DELETE FROM sql_attribute WHERE class_oid = ? +` + +func (q *Queries) DeleteAttributesByClass(ctx context.Context, classOid int64) error { + _, err := q.db.ExecContext(ctx, deleteAttributesByClass, classOid) + return err +} + +const deleteClass = `-- name: DeleteClass :exec +DELETE FROM sql_class WHERE oid = ? +` + +func (q *Queries) DeleteClass(ctx context.Context, oid int64) error { + _, err := q.db.ExecContext(ctx, deleteClass, oid) + return err +} + +const dialectFlag = `-- name: DialectFlag :one +SELECT value FROM sql_dialect_flag WHERE dialect_oid = ? AND key = ? +` + +type DialectFlagParams struct { + DialectOid int64 + Key string +} + +func (q *Queries) DialectFlag(ctx context.Context, arg DialectFlagParams) (string, error) { + row := q.db.QueryRowContext(ctx, dialectFlag, arg.DialectOid, arg.Key) + var value string + err := row.Scan(&value) + return value, err +} + +const dialectOID = `-- name: DialectOID :one +SELECT oid FROM sql_dialect WHERE name = ? +` + +func (q *Queries) DialectOID(ctx context.Context, name string) (int64, error) { + row := q.db.QueryRowContext(ctx, dialectOID, name) + var oid int64 + err := row.Scan(&oid) + return oid, err +} + +const findCast = `-- name: FindCast :one +SELECT source_type_oid, target_type_oid, proc_oid, context, dialect_oid +FROM sql_cast +WHERE source_type_oid = ? AND target_type_oid = ? +` + +type FindCastParams struct { + SourceTypeOid int64 + TargetTypeOid int64 +} + +func (q *Queries) FindCast(ctx context.Context, arg FindCastParams) (SqlCast, error) { + row := q.db.QueryRowContext(ctx, findCast, arg.SourceTypeOid, arg.TargetTypeOid) + var i SqlCast + err := row.Scan( + &i.SourceTypeOid, + &i.TargetTypeOid, + &i.ProcOid, + &i.Context, + &i.DialectOid, + ) + return i, err +} + +const findOperators = `-- name: FindOperators :many +SELECT oid, name, left_type_oid, right_type_oid, result_type_oid, proc_oid +FROM sql_operator +WHERE name = ?1 + AND (?2 = 0 OR left_type_oid = ?2) + AND (?3 = 0 OR right_type_oid = ?3) +` + +type FindOperatorsParams struct { + Name string + LeftTypeOid interface{} + RightTypeOid interface{} +} + +type FindOperatorsRow struct { + Oid int64 + Name string + LeftTypeOid sql.NullInt64 + RightTypeOid sql.NullInt64 + ResultTypeOid int64 + ProcOid sql.NullInt64 +} + +func (q *Queries) FindOperators(ctx context.Context, arg FindOperatorsParams) ([]FindOperatorsRow, error) { + rows, err := q.db.QueryContext(ctx, findOperators, arg.Name, arg.LeftTypeOid, arg.RightTypeOid) + if err != nil { + return nil, err + } + defer rows.Close() + var items []FindOperatorsRow + for rows.Next() { + var i FindOperatorsRow + if err := rows.Scan( + &i.Oid, + &i.Name, + &i.LeftTypeOid, + &i.RightTypeOid, + &i.ResultTypeOid, + &i.ProcOid, + ); err != nil { + return nil, err + } + items = append(items, i) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const findProcsAnyNamespace = `-- name: FindProcsAnyNamespace :many +SELECT oid, name, kind, return_type_oid, return_nullable +FROM sql_proc +WHERE name = ? +` + +type FindProcsAnyNamespaceRow struct { + Oid int64 + Name string + Kind string + ReturnTypeOid int64 + ReturnNullable int64 +} + +func (q *Queries) FindProcsAnyNamespace(ctx context.Context, name string) ([]FindProcsAnyNamespaceRow, error) { + rows, err := q.db.QueryContext(ctx, findProcsAnyNamespace, name) + if err != nil { + return nil, err + } + defer rows.Close() + var items []FindProcsAnyNamespaceRow + for rows.Next() { + var i FindProcsAnyNamespaceRow + if err := rows.Scan( + &i.Oid, + &i.Name, + &i.Kind, + &i.ReturnTypeOid, + &i.ReturnNullable, + ); err != nil { + return nil, err + } + items = append(items, i) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const findProcsInNamespaces = `-- name: FindProcsInNamespaces :many +SELECT oid, name, kind, return_type_oid, return_nullable +FROM sql_proc +WHERE name = ?1 + AND namespace_oid IN (/*SLICE:namespace_oids*/?) +` + +type FindProcsInNamespacesParams struct { + Name string + NamespaceOids []sql.NullInt64 +} + +type FindProcsInNamespacesRow struct { + Oid int64 + Name string + Kind string + ReturnTypeOid int64 + ReturnNullable int64 +} + +func (q *Queries) FindProcsInNamespaces(ctx context.Context, arg FindProcsInNamespacesParams) ([]FindProcsInNamespacesRow, error) { + query := findProcsInNamespaces + var queryParams []interface{} + queryParams = append(queryParams, arg.Name) + if len(arg.NamespaceOids) > 0 { + for _, v := range arg.NamespaceOids { + queryParams = append(queryParams, v) + } + query = strings.Replace(query, "/*SLICE:namespace_oids*/?", strings.Repeat(",?", len(arg.NamespaceOids))[1:], 1) + } else { + query = strings.Replace(query, "/*SLICE:namespace_oids*/?", "NULL", 1) + } + rows, err := q.db.QueryContext(ctx, query, queryParams...) + if err != nil { + return nil, err + } + defer rows.Close() + var items []FindProcsInNamespacesRow + for rows.Next() { + var i FindProcsInNamespacesRow + if err := rows.Scan( + &i.Oid, + &i.Name, + &i.Kind, + &i.ReturnTypeOid, + &i.ReturnNullable, + ); err != nil { + return nil, err + } + items = append(items, i) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const listClassColumns = `-- name: ListClassColumns :many +SELECT a.name AS column_name, t.name AS type_name, a.not_null +FROM sql_attribute a +JOIN sql_type t ON t.oid = a.type_oid +WHERE a.class_oid = ? +ORDER BY a.num +` + +type ListClassColumnsRow struct { + ColumnName string + TypeName string + NotNull int64 +} + +func (q *Queries) ListClassColumns(ctx context.Context, classOid int64) ([]ListClassColumnsRow, error) { + rows, err := q.db.QueryContext(ctx, listClassColumns, classOid) + if err != nil { + return nil, err + } + defer rows.Close() + var items []ListClassColumnsRow + for rows.Next() { + var i ListClassColumnsRow + if err := rows.Scan(&i.ColumnName, &i.TypeName, &i.NotNull); err != nil { + return nil, err + } + items = append(items, i) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const listNamespaces = `-- name: ListNamespaces :many +SELECT oid, name FROM sql_namespace ORDER BY oid +` + +func (q *Queries) ListNamespaces(ctx context.Context) ([]SqlNamespace, error) { + rows, err := q.db.QueryContext(ctx, listNamespaces) + if err != nil { + return nil, err + } + defer rows.Close() + var items []SqlNamespace + for rows.Next() { + var i SqlNamespace + if err := rows.Scan(&i.Oid, &i.Name); err != nil { + return nil, err + } + items = append(items, i) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const listTablesInNamespace = `-- name: ListTablesInNamespace :many +SELECT oid, name FROM sql_class +WHERE namespace_oid = ? AND kind = 'r' +ORDER BY oid +` + +type ListTablesInNamespaceRow struct { + Oid int64 + Name string +} + +func (q *Queries) ListTablesInNamespace(ctx context.Context, namespaceOid int64) ([]ListTablesInNamespaceRow, error) { + rows, err := q.db.QueryContext(ctx, listTablesInNamespace, namespaceOid) + if err != nil { + return nil, err + } + defer rows.Close() + var items []ListTablesInNamespaceRow + for rows.Next() { + var i ListTablesInNamespaceRow + if err := rows.Scan(&i.Oid, &i.Name); err != nil { + return nil, err + } + items = append(items, i) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const lookupAttribute = `-- name: LookupAttribute :one +SELECT ns.name AS schema_name, cls.name AS table_name, a.name AS column_name, a.num, + a.decl_type, a.type_length, a.type_scale, + a.auto_increment, a.is_primary_key, a.is_unique, a.not_null +FROM sql_attribute a +JOIN sql_class cls ON cls.oid = a.class_oid +JOIN sql_namespace ns ON ns.oid = cls.namespace_oid +WHERE a.oid = ? +` + +type LookupAttributeRow struct { + SchemaName string + TableName string + ColumnName string + Num int64 + DeclType string + TypeLength int64 + TypeScale int64 + AutoIncrement int64 + IsPrimaryKey int64 + IsUnique int64 + NotNull int64 +} + +func (q *Queries) LookupAttribute(ctx context.Context, oid int64) (LookupAttributeRow, error) { + row := q.db.QueryRowContext(ctx, lookupAttribute, oid) + var i LookupAttributeRow + err := row.Scan( + &i.SchemaName, + &i.TableName, + &i.ColumnName, + &i.Num, + &i.DeclType, + &i.TypeLength, + &i.TypeScale, + &i.AutoIncrement, + &i.IsPrimaryKey, + &i.IsUnique, + &i.NotNull, + ) + return i, err +} + +const lookupType = `-- name: LookupType :one +SELECT oid, name, category, typtype, preferred +FROM sql_type +WHERE oid = ? +` + +type LookupTypeRow struct { + Oid int64 + Name string + Category sql.NullString + Typtype string + Preferred int64 +} + +func (q *Queries) LookupType(ctx context.Context, oid int64) (LookupTypeRow, error) { + row := q.db.QueryRowContext(ctx, lookupType, oid) + var i LookupTypeRow + err := row.Scan( + &i.Oid, + &i.Name, + &i.Category, + &i.Typtype, + &i.Preferred, + ) + return i, err +} + +const namespaceOID = `-- name: NamespaceOID :one +SELECT oid FROM sql_namespace WHERE name = ? +` + +func (q *Queries) NamespaceOID(ctx context.Context, name string) (int64, error) { + row := q.db.QueryRowContext(ctx, namespaceOID, name) + var oid int64 + err := row.Scan(&oid) + return oid, err +} + +const procArgTypes = `-- name: ProcArgTypes :many +SELECT type_oid FROM sql_proc_arg +WHERE proc_oid = ? AND mode IN ('i', 'b', 'v') +ORDER BY ord +` + +func (q *Queries) ProcArgTypes(ctx context.Context, procOid int64) ([]int64, error) { + rows, err := q.db.QueryContext(ctx, procArgTypes, procOid) + if err != nil { + return nil, err + } + defer rows.Close() + var items []int64 + for rows.Next() { + var type_oid int64 + if err := rows.Scan(&type_oid); err != nil { + return nil, err + } + items = append(items, type_oid) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const resolveColumn = `-- name: ResolveColumn :one +SELECT a.oid, a.class_oid, a.name, t.name AS type_name, a.type_oid, a.not_null, + a.decl_type, a.type_length, a.type_scale, + a.auto_increment, a.is_primary_key, a.is_unique +FROM sql_attribute a +JOIN sql_class c ON c.oid = a.class_oid +JOIN sql_type t ON t.oid = a.type_oid +WHERE c.name = ?1 AND a.name = ?2 +` + +type ResolveColumnParams struct { + TableName string + ColumnName string +} + +type ResolveColumnRow struct { + Oid int64 + ClassOid int64 + Name string + TypeName string + TypeOid int64 + NotNull int64 + DeclType string + TypeLength int64 + TypeScale int64 + AutoIncrement int64 + IsPrimaryKey int64 + IsUnique int64 +} + +func (q *Queries) ResolveColumn(ctx context.Context, arg ResolveColumnParams) (ResolveColumnRow, error) { + row := q.db.QueryRowContext(ctx, resolveColumn, arg.TableName, arg.ColumnName) + var i ResolveColumnRow + err := row.Scan( + &i.Oid, + &i.ClassOid, + &i.Name, + &i.TypeName, + &i.TypeOid, + &i.NotNull, + &i.DeclType, + &i.TypeLength, + &i.TypeScale, + &i.AutoIncrement, + &i.IsPrimaryKey, + &i.IsUnique, + ) + return i, err +} + +const setAttributePrimaryKey = `-- name: SetAttributePrimaryKey :exec +UPDATE sql_attribute SET is_primary_key = 1, not_null = 1 +WHERE class_oid = ? AND name = ? +` + +type SetAttributePrimaryKeyParams struct { + ClassOid int64 + Name string +} + +func (q *Queries) SetAttributePrimaryKey(ctx context.Context, arg SetAttributePrimaryKeyParams) error { + _, err := q.db.ExecContext(ctx, setAttributePrimaryKey, arg.ClassOid, arg.Name) + return err +} + +const setAttributeUnique = `-- name: SetAttributeUnique :exec +UPDATE sql_attribute SET is_unique = 1 +WHERE class_oid = ? AND name = ? +` + +type SetAttributeUniqueParams struct { + ClassOid int64 + Name string +} + +func (q *Queries) SetAttributeUnique(ctx context.Context, arg SetAttributeUniqueParams) error { + _, err := q.db.ExecContext(ctx, setAttributeUnique, arg.ClassOid, arg.Name) + return err +} + +const setDialectFlag = `-- name: SetDialectFlag :exec +INSERT INTO sql_dialect_flag (dialect_oid, key, value) +VALUES (?, ?, ?) +ON CONFLICT(dialect_oid, key) DO UPDATE SET value = excluded.value +` + +type SetDialectFlagParams struct { + DialectOid int64 + Key string + Value string +} + +func (q *Queries) SetDialectFlag(ctx context.Context, arg SetDialectFlagParams) error { + _, err := q.db.ExecContext(ctx, setDialectFlag, arg.DialectOid, arg.Key, arg.Value) + return err +} + +const tableColumns = `-- name: TableColumns :many +SELECT a.oid, a.class_oid, a.name, t.name AS type_name, a.type_oid, a.not_null, + a.decl_type, a.type_length, a.type_scale, + a.auto_increment, a.is_primary_key, a.is_unique +FROM sql_attribute a +JOIN sql_class c ON c.oid = a.class_oid +JOIN sql_type t ON t.oid = a.type_oid +WHERE c.name = ? +ORDER BY a.num +` + +type TableColumnsRow struct { + Oid int64 + ClassOid int64 + Name string + TypeName string + TypeOid int64 + NotNull int64 + DeclType string + TypeLength int64 + TypeScale int64 + AutoIncrement int64 + IsPrimaryKey int64 + IsUnique int64 +} + +func (q *Queries) TableColumns(ctx context.Context, name string) ([]TableColumnsRow, error) { + rows, err := q.db.QueryContext(ctx, tableColumns, name) + if err != nil { + return nil, err + } + defer rows.Close() + var items []TableColumnsRow + for rows.Next() { + var i TableColumnsRow + if err := rows.Scan( + &i.Oid, + &i.ClassOid, + &i.Name, + &i.TypeName, + &i.TypeOid, + &i.NotNull, + &i.DeclType, + &i.TypeLength, + &i.TypeScale, + &i.AutoIncrement, + &i.IsPrimaryKey, + &i.IsUnique, + ); err != nil { + return nil, err + } + items = append(items, i) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const typeNameByOID = `-- name: TypeNameByOID :one +SELECT name FROM sql_type WHERE oid = ? +` + +func (q *Queries) TypeNameByOID(ctx context.Context, oid int64) (string, error) { + row := q.db.QueryRowContext(ctx, typeNameByOID, oid) + var name string + err := row.Scan(&name) + return name, err +} + +const typeOIDByName = `-- name: TypeOIDByName :one +SELECT t.oid +FROM sql_type t +JOIN sql_namespace ns ON ns.oid = t.namespace_oid +WHERE t.name = ?1 +ORDER BY + CASE ns.name + WHEN 'pg_catalog' THEN 0 + WHEN 'public' THEN 1 + ELSE 2 + END, + ns.name +LIMIT 1 +` + +func (q *Queries) TypeOIDByName(ctx context.Context, name string) (int64, error) { + row := q.db.QueryRowContext(ctx, typeOIDByName, name) + var oid int64 + err := row.Scan(&oid) + return oid, err +} diff --git a/internal/core/catalogdef/embed.go b/internal/core/catalogdef/embed.go new file mode 100644 index 0000000000..416961f2f2 --- /dev/null +++ b/internal/core/catalogdef/embed.go @@ -0,0 +1,6 @@ +package catalogdef + +import _ "embed" + +//go:embed schema.sql +var Schema string diff --git a/internal/core/catalogdef/query.sql b/internal/core/catalogdef/query.sql new file mode 100644 index 0000000000..b1e487d370 --- /dev/null +++ b/internal/core/catalogdef/query.sql @@ -0,0 +1,197 @@ +-- Queries against sqlc's own sql_* catalog tables, compiled by sqlc's +-- SQLite engine. Regenerate with `go generate ./internal/core/...`. + +-- ============================ sql_namespace ============================ + +-- name: CreateNamespace :execlastid +INSERT INTO sql_namespace (name) VALUES (?); + +-- name: NamespaceOID :one +SELECT oid FROM sql_namespace WHERE name = ?; + +-- name: ListNamespaces :many +SELECT oid, name FROM sql_namespace ORDER BY oid; + +-- ============================== sql_dialect ============================ + +-- name: CreateDialect :execlastid +INSERT INTO sql_dialect (name) VALUES (?); + +-- name: DialectOID :one +SELECT oid FROM sql_dialect WHERE name = ?; + +-- name: SetDialectFlag :exec +INSERT INTO sql_dialect_flag (dialect_oid, key, value) +VALUES (?, ?, ?) +ON CONFLICT(dialect_oid, key) DO UPDATE SET value = excluded.value; + +-- name: DialectFlag :one +SELECT value FROM sql_dialect_flag WHERE dialect_oid = ? AND key = ?; + +-- =============================== sql_type ============================== + +-- name: CreateType :execlastid +INSERT INTO sql_type + (name, size, typtype, category, preferred, namespace_oid, dialect_oid, element_oid) +VALUES (?, ?, ?, ?, ?, ?, ?, ?); + +-- name: TypeOIDByName :one +SELECT t.oid +FROM sql_type t +JOIN sql_namespace ns ON ns.oid = t.namespace_oid +WHERE t.name = sqlc.arg(name) +ORDER BY + CASE ns.name + WHEN 'pg_catalog' THEN 0 + WHEN 'public' THEN 1 + ELSE 2 + END, + ns.name +LIMIT 1; + +-- name: TypeNameByOID :one +SELECT name FROM sql_type WHERE oid = ?; + +-- name: LookupType :one +SELECT oid, name, category, typtype, preferred +FROM sql_type +WHERE oid = ?; + +-- =============================== sql_class ============================= + +-- name: CreateClass :execlastid +INSERT INTO sql_class (namespace_oid, name, kind) VALUES (?, ?, ?); + +-- name: ClassOID :one +SELECT oid FROM sql_class WHERE namespace_oid = ? AND name = ?; + +-- name: ClassOIDByName :one +SELECT oid FROM sql_class WHERE name = ? LIMIT 1; + +-- name: ListTablesInNamespace :many +SELECT oid, name FROM sql_class +WHERE namespace_oid = ? AND kind = 'r' +ORDER BY oid; + +-- name: DeleteClass :exec +DELETE FROM sql_class WHERE oid = ?; + +-- ============================= sql_attribute =========================== + +-- name: CreateAttribute :exec +INSERT INTO sql_attribute ( + class_oid, name, type_oid, not_null, has_default, num, + decl_type, type_length, type_scale, + auto_increment, is_primary_key, is_unique +) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?); + +-- name: SetAttributePrimaryKey :exec +UPDATE sql_attribute SET is_primary_key = 1, not_null = 1 +WHERE class_oid = ? AND name = ?; + +-- name: SetAttributeUnique :exec +UPDATE sql_attribute SET is_unique = 1 +WHERE class_oid = ? AND name = ?; + +-- name: DeleteAttributesByClass :exec +DELETE FROM sql_attribute WHERE class_oid = ?; + +-- name: ResolveColumn :one +SELECT a.oid, a.class_oid, a.name, t.name AS type_name, a.type_oid, a.not_null, + a.decl_type, a.type_length, a.type_scale, + a.auto_increment, a.is_primary_key, a.is_unique +FROM sql_attribute a +JOIN sql_class c ON c.oid = a.class_oid +JOIN sql_type t ON t.oid = a.type_oid +WHERE c.name = sqlc.arg(table_name) AND a.name = sqlc.arg(column_name); + +-- name: TableColumns :many +SELECT a.oid, a.class_oid, a.name, t.name AS type_name, a.type_oid, a.not_null, + a.decl_type, a.type_length, a.type_scale, + a.auto_increment, a.is_primary_key, a.is_unique +FROM sql_attribute a +JOIN sql_class c ON c.oid = a.class_oid +JOIN sql_type t ON t.oid = a.type_oid +WHERE c.name = ? +ORDER BY a.num; + +-- name: ClassAttributes :many +SELECT oid, name, type_oid, not_null +FROM sql_attribute +WHERE class_oid = ? +ORDER BY num; + +-- name: ListClassColumns :many +SELECT a.name AS column_name, t.name AS type_name, a.not_null +FROM sql_attribute a +JOIN sql_type t ON t.oid = a.type_oid +WHERE a.class_oid = ? +ORDER BY a.num; + +-- name: LookupAttribute :one +SELECT ns.name AS schema_name, cls.name AS table_name, a.name AS column_name, a.num, + a.decl_type, a.type_length, a.type_scale, + a.auto_increment, a.is_primary_key, a.is_unique, a.not_null +FROM sql_attribute a +JOIN sql_class cls ON cls.oid = a.class_oid +JOIN sql_namespace ns ON ns.oid = cls.namespace_oid +WHERE a.oid = ?; + +-- ============================ sql_constraint =========================== + +-- name: CreateConstraint :exec +INSERT INTO sql_constraint (class_oid, name, kind, columns) VALUES (?, ?, ?, ?); + +-- =============================== sql_proc ============================== + +-- name: CreateProc :execlastid +INSERT INTO sql_proc + (namespace_oid, dialect_oid, name, kind, + return_type_oid, return_set, return_nullable, strict, variadic_kind) +VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?); + +-- name: CreateProcArg :exec +INSERT INTO sql_proc_arg (proc_oid, ord, name, type_oid, mode, has_default) +VALUES (?, ?, ?, ?, ?, ?); + +-- name: ProcArgTypes :many +SELECT type_oid FROM sql_proc_arg +WHERE proc_oid = ? AND mode IN ('i', 'b', 'v') +ORDER BY ord; + +-- name: FindProcsAnyNamespace :many +SELECT oid, name, kind, return_type_oid, return_nullable +FROM sql_proc +WHERE name = ?; + +-- name: FindProcsInNamespaces :many +SELECT oid, name, kind, return_type_oid, return_nullable +FROM sql_proc +WHERE name = sqlc.arg(name) + AND namespace_oid IN (sqlc.slice(namespace_oids)); + +-- ============================= sql_operator ============================ + +-- name: CreateOperator :execlastid +INSERT INTO sql_operator + (namespace_oid, dialect_oid, name, + left_type_oid, right_type_oid, result_type_oid, proc_oid) +VALUES (?, ?, ?, ?, ?, ?, ?); + +-- name: FindOperators :many +SELECT oid, name, left_type_oid, right_type_oid, result_type_oid, proc_oid +FROM sql_operator +WHERE name = sqlc.arg(name) + AND (sqlc.arg(left_type_oid) = 0 OR left_type_oid = sqlc.arg(left_type_oid)) + AND (sqlc.arg(right_type_oid) = 0 OR right_type_oid = sqlc.arg(right_type_oid)); + +-- =============================== sql_cast ============================== + +-- name: CreateCast :exec +INSERT INTO sql_cast (source_type_oid, target_type_oid, proc_oid, context, dialect_oid) +VALUES (?, ?, ?, ?, ?); + +-- name: FindCast :one +SELECT source_type_oid, target_type_oid, proc_oid, context, dialect_oid +FROM sql_cast +WHERE source_type_oid = ? AND target_type_oid = ?; diff --git a/internal/core/catalogdef/schema.sql b/internal/core/catalogdef/schema.sql new file mode 100644 index 0000000000..e9d9741d4b --- /dev/null +++ b/internal/core/catalogdef/schema.sql @@ -0,0 +1,156 @@ +-- sql_namespace: schemas / namespaces +CREATE TABLE sql_namespace ( + oid INTEGER PRIMARY KEY AUTOINCREMENT, + name TEXT NOT NULL UNIQUE +); + +-- sql_dialect: registered SQL dialects (postgresql, sqlite, mysql, ...). +CREATE TABLE sql_dialect ( + oid INTEGER PRIMARY KEY AUTOINCREMENT, + name TEXT NOT NULL UNIQUE +); + +-- sql_dialect_flag: per-dialect configuration knobs (case-folding, +-- identifier quoting, alias scoping rules, etc.). Values are opaque +-- strings; the analyzer interprets them per key. +CREATE TABLE sql_dialect_flag ( + dialect_oid INTEGER NOT NULL REFERENCES sql_dialect(oid), + key TEXT NOT NULL, + value TEXT NOT NULL, + PRIMARY KEY (dialect_oid, key) +); + +-- sql_type: data types. Modeled on pg_type. +-- typtype: 'b'ase | 'c'omposite | 'd'omain | 'e'num | 'p'seudo | 'r'ange +-- category: 'N'umeric | 'S'tring | 'B'oolean | 'D'atetime | 'A'rray | +-- 'C'omposite | 'E'num | 'U'serdef | 'X'unknown +-- preferred: tie-breaker for implicit cast resolution within a category +-- element_oid: for arrays, points at the element type +-- dialect_oid: NULL = standard / shared across dialects +CREATE TABLE sql_type ( + oid INTEGER PRIMARY KEY AUTOINCREMENT, + namespace_oid INTEGER NOT NULL REFERENCES sql_namespace(oid), + dialect_oid INTEGER REFERENCES sql_dialect(oid), + name TEXT NOT NULL, + size INTEGER NOT NULL DEFAULT 0, + typtype TEXT NOT NULL DEFAULT 'b', + category TEXT, + preferred INTEGER NOT NULL DEFAULT 0, + element_oid INTEGER REFERENCES sql_type(oid), + UNIQUE (namespace_oid, name) +); +CREATE INDEX idx_sql_type_name ON sql_type(name); + +-- sql_class: relations (tables, views, indexes). +-- kind: 'r' = table, 'v' = view, 'i' = index, 'c' = composite type, 'f' = foreign +CREATE TABLE sql_class ( + oid INTEGER PRIMARY KEY AUTOINCREMENT, + namespace_oid INTEGER NOT NULL REFERENCES sql_namespace(oid), + name TEXT NOT NULL, + kind TEXT NOT NULL DEFAULT 'r', + UNIQUE(namespace_oid, name) +); + +-- sql_attribute: columns of a relation. +-- decl_type: original declared type string before normalization +-- (e.g. VARCHAR(10), BIGINT UNSIGNED, INTEGER PRIMARY KEY). +-- Useful for SQLite where multiple syntaxes collapse to +-- one of five affinities, and as a debugging aid. +-- type_length: length / precision (varchar(N), numeric(p,s).p, +-- char(N), bit(N)). 0 = unspecified. +-- type_scale: scale for numeric/decimal. 0 = unspecified. +-- auto_increment: rowid alias (sqlite INTEGER PRIMARY KEY), AUTOINCREMENT, +-- pg serial/bigserial/identity, mysql AUTO_INCREMENT. +-- is_primary_key: this column participates in the relation's primary key. +-- Set both for inline-column PK and for table-level PK. +-- is_unique: column has a UNIQUE constraint or a single-column UNIQUE +-- table constraint. +CREATE TABLE sql_attribute ( + oid INTEGER PRIMARY KEY AUTOINCREMENT, + class_oid INTEGER NOT NULL REFERENCES sql_class(oid), + name TEXT NOT NULL, + type_oid INTEGER NOT NULL REFERENCES sql_type(oid), + not_null INTEGER NOT NULL DEFAULT 0, + has_default INTEGER NOT NULL DEFAULT 0, + num INTEGER NOT NULL, -- ordinal position (1-based) + decl_type TEXT NOT NULL DEFAULT '', + type_length INTEGER NOT NULL DEFAULT 0, + type_scale INTEGER NOT NULL DEFAULT 0, + auto_increment INTEGER NOT NULL DEFAULT 0, + is_primary_key INTEGER NOT NULL DEFAULT 0, + is_unique INTEGER NOT NULL DEFAULT 0, + UNIQUE(class_oid, name), + UNIQUE(class_oid, num) +); + +-- sql_constraint: constraints on a relation. +-- kind: 'p' = primary key, 'f' = foreign key, 'u' = unique, 'c' = check +CREATE TABLE sql_constraint ( + oid INTEGER PRIMARY KEY AUTOINCREMENT, + class_oid INTEGER NOT NULL REFERENCES sql_class(oid), + name TEXT NOT NULL DEFAULT '', + kind TEXT NOT NULL, + columns TEXT NOT NULL DEFAULT '' -- comma-separated attribute nums +); + +-- sql_proc: functions, aggregates, window functions, procedures. +-- Modeled on pg_proc. +-- kind: 'f' = function, 'a' = aggregate, 'w' = window, 'p' = procedure +-- variadic_kind: 'n' = none, 'a' = array (VARIADIC any[]), 'v' = variadic-any +-- return_set: 1 if SETOF / table-returning +CREATE TABLE sql_proc ( + oid INTEGER PRIMARY KEY AUTOINCREMENT, + namespace_oid INTEGER REFERENCES sql_namespace(oid), + dialect_oid INTEGER REFERENCES sql_dialect(oid), + name TEXT NOT NULL, + kind TEXT NOT NULL DEFAULT 'f', + return_type_oid INTEGER NOT NULL REFERENCES sql_type(oid), + return_set INTEGER NOT NULL DEFAULT 0, + return_nullable INTEGER NOT NULL DEFAULT 1, + strict INTEGER NOT NULL DEFAULT 0, + variadic_kind TEXT NOT NULL DEFAULT 'n' +); + +-- sql_proc_arg: ordered argument list for a proc. +-- mode: 'i' = in, 'o' = out, 'b' = both, 't' = table, 'v' = variadic +CREATE TABLE sql_proc_arg ( + proc_oid INTEGER NOT NULL REFERENCES sql_proc(oid), + ord INTEGER NOT NULL, + name TEXT NOT NULL DEFAULT '', + type_oid INTEGER NOT NULL REFERENCES sql_type(oid), + mode TEXT NOT NULL DEFAULT 'i', + has_default INTEGER NOT NULL DEFAULT 0, + PRIMARY KEY (proc_oid, ord) +); + +-- sql_operator: operator overloads. +-- left_type_oid is NULL for prefix unary; right_type_oid is NULL for postfix. +CREATE TABLE sql_operator ( + oid INTEGER PRIMARY KEY AUTOINCREMENT, + namespace_oid INTEGER REFERENCES sql_namespace(oid), + dialect_oid INTEGER REFERENCES sql_dialect(oid), + name TEXT NOT NULL, + left_type_oid INTEGER REFERENCES sql_type(oid), + right_type_oid INTEGER REFERENCES sql_type(oid), + result_type_oid INTEGER NOT NULL REFERENCES sql_type(oid), + proc_oid INTEGER REFERENCES sql_proc(oid), + commutator_oid INTEGER REFERENCES sql_operator(oid), + negator_oid INTEGER REFERENCES sql_operator(oid) +); + +-- sql_cast: type coercion rules. +-- context: 'i' = implicit, 'a' = assignment-only, 'e' = explicit-only +-- proc_oid NULL = binary-coercible (no function needed) +CREATE TABLE sql_cast ( + source_type_oid INTEGER NOT NULL REFERENCES sql_type(oid), + target_type_oid INTEGER NOT NULL REFERENCES sql_type(oid), + proc_oid INTEGER REFERENCES sql_proc(oid), + context TEXT NOT NULL DEFAULT 'e', + dialect_oid INTEGER REFERENCES sql_dialect(oid), + PRIMARY KEY (source_type_oid, target_type_oid) +); + +-- Resolution-speed indexes. +CREATE INDEX idx_sql_proc_name ON sql_proc(name, namespace_oid); +CREATE INDEX idx_sql_operator_name ON sql_operator(name, left_type_oid, right_type_oid); +CREATE INDEX idx_sql_attribute_name ON sql_attribute(name); diff --git a/internal/core/class.go b/internal/core/class.go new file mode 100644 index 0000000000..33e5c69192 --- /dev/null +++ b/internal/core/class.go @@ -0,0 +1,67 @@ +package core + +import ( + "context" + "fmt" + + "github.com/sqlc-dev/sqlc/internal/core/catalogdb" +) + +func (c *Catalog) CreateClass(namespaceOID int64, name string, kind string) (int64, error) { + oid, err := c.q.CreateClass(context.Background(), catalogdb.CreateClassParams{ + NamespaceOid: namespaceOID, + Name: name, + Kind: kind, + }) + if err != nil { + return 0, fmt.Errorf("create class %q: %w", name, err) + } + return oid, nil +} + +func (c *Catalog) ClassOID(namespaceOID int64, name string) (int64, error) { + oid, err := c.q.ClassOID(context.Background(), catalogdb.ClassOIDParams{ + NamespaceOid: namespaceOID, + Name: name, + }) + if err != nil { + return 0, fmt.Errorf("class %q: %w", name, err) + } + return oid, nil +} + +func (c *Catalog) ClassOIDByName(name string) (int64, error) { + oid, err := c.q.ClassOIDByName(context.Background(), name) + if err != nil { + return 0, fmt.Errorf("class %q: %w", name, err) + } + return oid, nil +} + +func (c *Catalog) DropClass(classOID int64) error { + ctx := context.Background() + if err := c.q.DeleteAttributesByClass(ctx, classOID); err != nil { + return fmt.Errorf("drop class %d attributes: %w", classOID, err) + } + if err := c.q.DeleteClass(ctx, classOID); err != nil { + return fmt.Errorf("drop class %d: %w", classOID, err) + } + return nil +} + +type ClassInfo struct { + OID int64 + Name string +} + +func (c *Catalog) TablesInNamespace(namespaceOID int64) ([]ClassInfo, error) { + rows, err := c.q.ListTablesInNamespace(context.Background(), namespaceOID) + if err != nil { + return nil, fmt.Errorf("list tables in namespace %d: %w", namespaceOID, err) + } + out := make([]ClassInfo, 0, len(rows)) + for _, r := range rows { + out = append(out, ClassInfo{OID: r.Oid, Name: r.Name}) + } + return out, nil +} diff --git a/internal/core/constraint.go b/internal/core/constraint.go new file mode 100644 index 0000000000..d75e3c2abf --- /dev/null +++ b/internal/core/constraint.go @@ -0,0 +1,21 @@ +package core + +import ( + "context" + "fmt" + + "github.com/sqlc-dev/sqlc/internal/core/catalogdb" +) + +func (c *Catalog) CreateConstraint(classOID int64, name string, kind string, columns string) error { + err := c.q.CreateConstraint(context.Background(), catalogdb.CreateConstraintParams{ + ClassOid: classOID, + Name: name, + Kind: kind, + Columns: columns, + }) + if err != nil { + return fmt.Errorf("create constraint %q on class %d: %w", name, classOID, err) + } + return nil +} diff --git a/internal/core/dialect.go b/internal/core/dialect.go new file mode 100644 index 0000000000..283c0324b4 --- /dev/null +++ b/internal/core/dialect.go @@ -0,0 +1,47 @@ +package core + +import ( + "context" + "fmt" + + "github.com/sqlc-dev/sqlc/internal/core/catalogdb" +) + +func (c *Catalog) CreateDialect(name string) (int64, error) { + oid, err := c.q.CreateDialect(context.Background(), name) + if err != nil { + return 0, fmt.Errorf("create dialect %q: %w", name, err) + } + return oid, nil +} + +func (c *Catalog) DialectOID(name string) (int64, error) { + oid, err := c.q.DialectOID(context.Background(), name) + if err != nil { + return 0, fmt.Errorf("dialect %q: %w", name, err) + } + return oid, nil +} + +func (c *Catalog) SetDialectFlag(dialectOID int64, key, value string) error { + err := c.q.SetDialectFlag(context.Background(), catalogdb.SetDialectFlagParams{ + DialectOid: dialectOID, + Key: key, + Value: value, + }) + if err != nil { + return fmt.Errorf("set dialect flag %s.%s: %w", key, value, err) + } + return nil +} + +func (c *Catalog) DialectFlag(dialectOID int64, key string) (string, error) { + value, err := c.q.DialectFlag(context.Background(), catalogdb.DialectFlagParams{ + DialectOid: dialectOID, + Key: key, + }) + if err != nil { + return "", nil + } + return value, nil +} diff --git a/internal/core/dialect_test.go b/internal/core/dialect_test.go new file mode 100644 index 0000000000..e830a035f4 --- /dev/null +++ b/internal/core/dialect_test.go @@ -0,0 +1,42 @@ +package core + +import "testing" + +func TestDialectAndFlags(t *testing.T) { + cat, err := New() + if err != nil { + t.Fatal(err) + } + defer cat.Close() + + pgOID, err := cat.CreateDialect("postgresql") + if err != nil { + t.Fatal(err) + } + + got, err := cat.DialectOID("postgresql") + if err != nil || got != pgOID { + t.Fatalf("DialectOID: got %d (err=%v), want %d", got, err, pgOID) + } + + if err := cat.SetDialectFlag(pgOID, "identifier_case", "fold_lower"); err != nil { + t.Fatal(err) + } + v, err := cat.DialectFlag(pgOID, "identifier_case") + if err != nil || v != "fold_lower" { + t.Errorf("DialectFlag: got %q (err=%v), want fold_lower", v, err) + } + + if err := cat.SetDialectFlag(pgOID, "identifier_case", "fold_upper"); err != nil { + t.Fatal(err) + } + v, _ = cat.DialectFlag(pgOID, "identifier_case") + if v != "fold_upper" { + t.Errorf("upsert: got %q, want fold_upper", v) + } + + v, err = cat.DialectFlag(pgOID, "nonexistent") + if err != nil || v != "" { + t.Errorf("missing flag: got %q (err=%v), want empty", v, err) + } +} diff --git a/internal/core/namespace.go b/internal/core/namespace.go new file mode 100644 index 0000000000..b59405336e --- /dev/null +++ b/internal/core/namespace.go @@ -0,0 +1,39 @@ +package core + +import ( + "context" + "fmt" +) + +func (c *Catalog) CreateNamespace(name string) (int64, error) { + oid, err := c.q.CreateNamespace(context.Background(), name) + if err != nil { + return 0, fmt.Errorf("create namespace %q: %w", name, err) + } + return oid, nil +} + +func (c *Catalog) NamespaceOID(name string) (int64, error) { + oid, err := c.q.NamespaceOID(context.Background(), name) + if err != nil { + return 0, fmt.Errorf("namespace %q: %w", name, err) + } + return oid, nil +} + +type NamespaceInfo struct { + OID int64 + Name string +} + +func (c *Catalog) Namespaces() ([]NamespaceInfo, error) { + rows, err := c.q.ListNamespaces(context.Background()) + if err != nil { + return nil, fmt.Errorf("list namespaces: %w", err) + } + out := make([]NamespaceInfo, 0, len(rows)) + for _, r := range rows { + out = append(out, NamespaceInfo{OID: r.Oid, Name: r.Name}) + } + return out, nil +} diff --git a/internal/core/operator.go b/internal/core/operator.go new file mode 100644 index 0000000000..7eae634aeb --- /dev/null +++ b/internal/core/operator.go @@ -0,0 +1,66 @@ +package core + +import ( + "context" + "fmt" + + "github.com/sqlc-dev/sqlc/internal/core/catalogdb" +) + +type OperatorSpec struct { + Name string + NamespaceOID int64 + DialectOID int64 + LeftTypeOID int64 + RightTypeOID int64 + ResultTypeOID int64 + ProcOID int64 +} + +func (c *Catalog) CreateOperator(o OperatorSpec) (int64, error) { + oid, err := c.q.CreateOperator(context.Background(), catalogdb.CreateOperatorParams{ + NamespaceOid: nullInt64(o.NamespaceOID), + DialectOid: nullInt64(o.DialectOID), + Name: o.Name, + LeftTypeOid: nullInt64(o.LeftTypeOID), + RightTypeOid: nullInt64(o.RightTypeOID), + ResultTypeOid: o.ResultTypeOID, + ProcOid: nullInt64(o.ProcOID), + }) + if err != nil { + return 0, fmt.Errorf("create operator %q: %w", o.Name, err) + } + return oid, nil +} + +type OperatorOverload struct { + OID int64 + Name string + LeftTypeOID int64 + RightTypeOID int64 + ResultTypeOID int64 + ProcOID int64 +} + +func (c *Catalog) FindOperators(name string, leftTypeOID, rightTypeOID int64) ([]OperatorOverload, error) { + rows, err := c.q.FindOperators(context.Background(), catalogdb.FindOperatorsParams{ + Name: name, + LeftTypeOid: leftTypeOID, + RightTypeOid: rightTypeOID, + }) + if err != nil { + return nil, fmt.Errorf("find operators %q: %w", name, err) + } + out := make([]OperatorOverload, 0, len(rows)) + for _, r := range rows { + out = append(out, OperatorOverload{ + OID: r.Oid, + Name: r.Name, + LeftTypeOID: orZero(r.LeftTypeOid), + RightTypeOID: orZero(r.RightTypeOid), + ResultTypeOID: r.ResultTypeOid, + ProcOID: orZero(r.ProcOid), + }) + } + return out, nil +} diff --git a/internal/core/operator_test.go b/internal/core/operator_test.go new file mode 100644 index 0000000000..1aa5ac045b --- /dev/null +++ b/internal/core/operator_test.go @@ -0,0 +1,46 @@ +package core + +import "testing" + +func TestOperatorCreateAndFind(t *testing.T) { + cat, err := New() + if err != nil { + t.Fatal(err) + } + defer cat.Close() + + intOID, _ := cat.CreateType("integer", 4) + boolOID, _ := cat.CreateType("boolean", 1) + + gtOID, err := cat.CreateOperator(OperatorSpec{ + Name: ">", + LeftTypeOID: intOID, + RightTypeOID: intOID, + ResultTypeOID: boolOID, + }) + if err != nil { + t.Fatal(err) + } + if gtOID == 0 { + t.Fatal("expected non-zero operator oid") + } + + got, err := cat.FindOperators(">", intOID, intOID) + if err != nil { + t.Fatal(err) + } + if len(got) != 1 { + t.Fatalf("FindOperators: want 1, got %d", len(got)) + } + if got[0].ResultTypeOID != boolOID { + t.Errorf("result type: got %d, want %d", got[0].ResultTypeOID, boolOID) + } + + all, err := cat.FindOperators(">", 0, 0) + if err != nil { + t.Fatal(err) + } + if len(all) != 1 { + t.Errorf("listing >: want 1, got %d", len(all)) + } +} diff --git a/internal/core/proc.go b/internal/core/proc.go new file mode 100644 index 0000000000..a726c9f459 --- /dev/null +++ b/internal/core/proc.go @@ -0,0 +1,137 @@ +package core + +import ( + "context" + "database/sql" + "fmt" + "strings" + + "github.com/sqlc-dev/sqlc/internal/core/catalogdb" +) + +type ProcSpec struct { + Name string + NamespaceOID int64 + DialectOID int64 + Kind string + ReturnTypeOID int64 + ReturnSet bool + ReturnNullable bool + Strict bool + VariadicKind string + Args []ProcArg +} + +type ProcArg struct { + Name string + TypeOID int64 + Mode string + HasDefault bool +} + +func (c *Catalog) CreateProc(p ProcSpec) (int64, error) { + if p.Kind == "" { + p.Kind = "f" + } + if p.VariadicKind == "" { + p.VariadicKind = "n" + } + ctx := context.Background() + procOID, err := c.q.CreateProc(ctx, catalogdb.CreateProcParams{ + NamespaceOid: nullInt64(p.NamespaceOID), + DialectOid: nullInt64(p.DialectOID), + Name: strings.ToLower(p.Name), + Kind: p.Kind, + ReturnTypeOid: p.ReturnTypeOID, + ReturnSet: boolToInt64(p.ReturnSet), + ReturnNullable: boolToInt64(p.ReturnNullable), + Strict: boolToInt64(p.Strict), + VariadicKind: p.VariadicKind, + }) + if err != nil { + return 0, fmt.Errorf("create proc %q: %w", p.Name, err) + } + for i, a := range p.Args { + mode := a.Mode + if mode == "" { + mode = "i" + } + err := c.q.CreateProcArg(ctx, catalogdb.CreateProcArgParams{ + ProcOid: procOID, + Ord: int64(i + 1), + Name: a.Name, + TypeOid: a.TypeOID, + Mode: mode, + HasDefault: boolToInt64(a.HasDefault), + }) + if err != nil { + return 0, fmt.Errorf("create proc %q arg %d: %w", p.Name, i+1, err) + } + } + return procOID, nil +} + +type ProcOverload struct { + OID int64 + Name string + Kind string + ReturnTypeOID int64 + ReturnNullable bool + ArgTypes []int64 +} + +func (c *Catalog) FindProcs(name string, namespaceOIDs []int64) ([]ProcOverload, error) { + ctx := context.Background() + lname := strings.ToLower(name) + + var out []ProcOverload + if len(namespaceOIDs) == 0 { + rows, err := c.q.FindProcsAnyNamespace(ctx, lname) + if err != nil { + return nil, fmt.Errorf("find procs %q: %w", name, err) + } + for _, r := range rows { + out = append(out, ProcOverload{ + OID: r.Oid, + Name: r.Name, + Kind: r.Kind, + ReturnTypeOID: r.ReturnTypeOid, + ReturnNullable: r.ReturnNullable != 0, + }) + } + } else { + nss := make([]sql.NullInt64, len(namespaceOIDs)) + for i, ns := range namespaceOIDs { + nss[i] = sql.NullInt64{Int64: ns, Valid: true} + } + rows, err := c.q.FindProcsInNamespaces(ctx, catalogdb.FindProcsInNamespacesParams{ + Name: lname, + NamespaceOids: nss, + }) + if err != nil { + return nil, fmt.Errorf("find procs %q: %w", name, err) + } + for _, r := range rows { + out = append(out, ProcOverload{ + OID: r.Oid, + Name: r.Name, + Kind: r.Kind, + ReturnTypeOID: r.ReturnTypeOid, + ReturnNullable: r.ReturnNullable != 0, + }) + } + } + + for i := range out { + argTypes, err := c.procArgTypes(out[i].OID) + if err != nil { + return nil, err + } + out[i].ArgTypes = argTypes + } + return out, nil +} + +func (c *Catalog) procArgTypes(procOID int64) ([]int64, error) { + return c.q.ProcArgTypes(context.Background(), procOID) +} diff --git a/internal/core/proc_test.go b/internal/core/proc_test.go new file mode 100644 index 0000000000..7b69b77956 --- /dev/null +++ b/internal/core/proc_test.go @@ -0,0 +1,67 @@ +package core + +import "testing" + +func TestProcCreateAndFind(t *testing.T) { + cat, err := New() + if err != nil { + t.Fatal(err) + } + defer cat.Close() + + intOID, err := cat.CreateType("integer", 4) + if err != nil { + t.Fatal(err) + } + textOID, err := cat.CreateType("text", 0) + if err != nil { + t.Fatal(err) + } + + procOID, err := cat.CreateProc(ProcSpec{ + Name: "length", + Kind: "f", + ReturnTypeOID: intOID, + ReturnNullable: true, + Args: []ProcArg{{Name: "s", TypeOID: textOID}}, + }) + if err != nil { + t.Fatalf("create length: %v", err) + } + if procOID == 0 { + t.Fatal("expected non-zero proc oid") + } + + if _, err := cat.CreateProc(ProcSpec{ + Name: "concat", + ReturnTypeOID: textOID, + Args: []ProcArg{ + {Name: "a", TypeOID: textOID}, + {Name: "b", TypeOID: textOID}, + }, + }); err != nil { + t.Fatal(err) + } + + got, err := cat.FindProcs("length", nil) + if err != nil { + t.Fatal(err) + } + if len(got) != 1 { + t.Fatalf("FindProcs length: want 1, got %d", len(got)) + } + if got[0].ReturnTypeOID != intOID { + t.Errorf("length return type: got %d, want %d", got[0].ReturnTypeOID, intOID) + } + if len(got[0].ArgTypes) != 1 || got[0].ArgTypes[0] != textOID { + t.Errorf("length args: got %v, want [%d]", got[0].ArgTypes, textOID) + } + + got, err = cat.FindProcs("concat", nil) + if err != nil { + t.Fatal(err) + } + if len(got) != 1 || len(got[0].ArgTypes) != 2 { + t.Fatalf("FindProcs concat: got %+v", got) + } +} diff --git a/internal/core/sqlc.json b/internal/core/sqlc.json new file mode 100644 index 0000000000..3b4763b697 --- /dev/null +++ b/internal/core/sqlc.json @@ -0,0 +1,16 @@ +{ + "version": "2", + "sql": [ + { + "engine": "sqlite", + "schema": "catalogdef/schema.sql", + "queries": "catalogdef/query.sql", + "gen": { + "go": { + "package": "catalogdb", + "out": "catalogdb" + } + } + } + ] +} diff --git a/internal/core/types.go b/internal/core/types.go new file mode 100644 index 0000000000..7515f95372 --- /dev/null +++ b/internal/core/types.go @@ -0,0 +1,133 @@ +package core + +import ( + "context" + "database/sql" + "fmt" + "strings" + + "github.com/sqlc-dev/sqlc/internal/core/catalogdb" +) + +type TypeSpec struct { + Name string + Size int + Typtype string + Category string + Preferred bool + NamespaceOID int64 + DialectOID int64 + ElementOID int64 +} + +func (c *Catalog) CreateType(name string, size int) (int64, error) { + return c.CreateTypeSpec(TypeSpec{Name: name, Size: size, Typtype: "b"}) +} + +func (c *Catalog) CreateTypeSpec(t TypeSpec) (int64, error) { + if t.Typtype == "" { + t.Typtype = "b" + } + if t.NamespaceOID == 0 { + oid, err := c.NamespaceOID("public") + if err != nil { + return 0, fmt.Errorf("create type %q: default namespace: %w", t.Name, err) + } + t.NamespaceOID = oid + } + oid, err := c.q.CreateType(context.Background(), catalogdb.CreateTypeParams{ + Name: strings.ToLower(t.Name), + Size: int64(t.Size), + Typtype: t.Typtype, + Category: nullString(t.Category), + Preferred: boolToInt64(t.Preferred), + NamespaceOid: t.NamespaceOID, + DialectOid: nullInt64(t.DialectOID), + ElementOid: nullInt64(t.ElementOID), + }) + if err != nil { + return 0, fmt.Errorf("create type %q: %w", t.Name, err) + } + return oid, nil +} + +func (c *Catalog) TypeOID(name string) (int64, error) { + oid, err := c.q.TypeOIDByName(context.Background(), strings.ToLower(name)) + if err != nil { + return 0, fmt.Errorf("type %q: %w", name, err) + } + return oid, nil +} + +func (c *Catalog) TypeName(oid int64) (string, error) { + name, err := c.q.TypeNameByOID(context.Background(), oid) + if err != nil { + return "", fmt.Errorf("type oid %d: %w", oid, err) + } + return name, nil +} + +type TypeInfo struct { + OID int64 + Name string + Category string + Typtype string + Preferred bool +} + +func (c *Catalog) LookupType(oid int64) (TypeInfo, error) { + row, err := c.q.LookupType(context.Background(), oid) + if err != nil { + return TypeInfo{}, fmt.Errorf("lookup type oid %d: %w", oid, err) + } + return TypeInfo{ + OID: row.Oid, + Name: row.Name, + Category: row.Category.String, + Typtype: row.Typtype, + Preferred: row.Preferred != 0, + }, nil +} + +func nullableOID(oid int64) any { + if oid == 0 { + return nil + } + return oid +} + +func nullableString(s string) any { + if s == "" { + return nil + } + return s +} + +func boolToInt(b bool) int { + if b { + return 1 + } + return 0 +} + +func nullInt64(oid int64) sql.NullInt64 { + return sql.NullInt64{Int64: oid, Valid: oid != 0} +} + +func nullString(s string) sql.NullString { + return sql.NullString{String: s, Valid: s != ""} +} + +func boolToInt64(b bool) int64 { + if b { + return 1 + } + return 0 +} + +func orZero(n sql.NullInt64) int64 { + if n.Valid { + return n.Int64 + } + return 0 +} diff --git a/internal/endtoend/testdata/clickhouse_select/clickhouse/db/db.go b/internal/endtoend/testdata/clickhouse_select/clickhouse/db/db.go new file mode 100644 index 0000000000..f43598b1eb --- /dev/null +++ b/internal/endtoend/testdata/clickhouse_select/clickhouse/db/db.go @@ -0,0 +1,31 @@ +// Code generated by sqlc. DO NOT EDIT. +// versions: +// sqlc v1.31.1 + +package db + +import ( + "context" + "database/sql" +) + +type DBTX interface { + ExecContext(context.Context, string, ...interface{}) (sql.Result, error) + PrepareContext(context.Context, string) (*sql.Stmt, error) + QueryContext(context.Context, string, ...interface{}) (*sql.Rows, error) + QueryRowContext(context.Context, string, ...interface{}) *sql.Row +} + +func New(db DBTX) *Queries { + return &Queries{db: db} +} + +type Queries struct { + db DBTX +} + +func (q *Queries) WithTx(tx *sql.Tx) *Queries { + return &Queries{ + db: tx, + } +} diff --git a/internal/endtoend/testdata/clickhouse_select/clickhouse/db/models.go b/internal/endtoend/testdata/clickhouse_select/clickhouse/db/models.go new file mode 100644 index 0000000000..21f493ecbc --- /dev/null +++ b/internal/endtoend/testdata/clickhouse_select/clickhouse/db/models.go @@ -0,0 +1,5 @@ +// Code generated by sqlc. DO NOT EDIT. +// versions: +// sqlc v1.31.1 + +package db diff --git a/internal/endtoend/testdata/clickhouse_select/clickhouse/db/query.sql.go b/internal/endtoend/testdata/clickhouse_select/clickhouse/db/query.sql.go new file mode 100644 index 0000000000..19233a0b45 --- /dev/null +++ b/internal/endtoend/testdata/clickhouse_select/clickhouse/db/query.sql.go @@ -0,0 +1,52 @@ +// Code generated by sqlc. DO NOT EDIT. +// versions: +// sqlc v1.31.1 +// source: query.sql + +package db + +import ( + "context" + "time" +) + +const listEvents = `-- name: ListEvents :many +SELECT id, name, tag, amount, created FROM events; +` + +type ListEventsRow struct { + ID uint64 + Name string + Tag *string + Amount float64 + Created time.Time +} + +func (q *Queries) ListEvents(ctx context.Context) ([]ListEventsRow, error) { + rows, err := q.db.QueryContext(ctx, listEvents) + if err != nil { + return nil, err + } + defer rows.Close() + var items []ListEventsRow + for rows.Next() { + var i ListEventsRow + if err := rows.Scan( + &i.ID, + &i.Name, + &i.Tag, + &i.Amount, + &i.Created, + ); err != nil { + return nil, err + } + items = append(items, i) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} diff --git a/internal/endtoend/testdata/clickhouse_select/clickhouse/query.sql b/internal/endtoend/testdata/clickhouse_select/clickhouse/query.sql new file mode 100644 index 0000000000..4c756ec0f8 --- /dev/null +++ b/internal/endtoend/testdata/clickhouse_select/clickhouse/query.sql @@ -0,0 +1,2 @@ +-- name: ListEvents :many +SELECT id, name, tag, amount, created FROM events; diff --git a/internal/endtoend/testdata/clickhouse_select/clickhouse/schema.sql b/internal/endtoend/testdata/clickhouse_select/clickhouse/schema.sql new file mode 100644 index 0000000000..29960ee63d --- /dev/null +++ b/internal/endtoend/testdata/clickhouse_select/clickhouse/schema.sql @@ -0,0 +1,7 @@ +CREATE TABLE events ( + id UInt64, + name String, + tag Nullable(String), + amount Float64, + created DateTime +) ENGINE = MergeTree ORDER BY id; diff --git a/internal/endtoend/testdata/clickhouse_select/clickhouse/sqlc.json b/internal/endtoend/testdata/clickhouse_select/clickhouse/sqlc.json new file mode 100644 index 0000000000..995102c9d6 --- /dev/null +++ b/internal/endtoend/testdata/clickhouse_select/clickhouse/sqlc.json @@ -0,0 +1,16 @@ +{ + "version": "2", + "sql": [ + { + "engine": "clickhouse", + "queries": "query.sql", + "schema": "schema.sql", + "gen": { + "go": { + "package": "db", + "out": "db" + } + } + } + ] +} diff --git a/internal/endtoend/testdata/parse_basic/clickhouse/stdout.txt b/internal/endtoend/testdata/parse_basic/clickhouse/stdout.txt index e2c49df3fa..28a5ce7e1f 100644 --- a/internal/endtoend/testdata/parse_basic/clickhouse/stdout.txt +++ b/internal/endtoend/testdata/parse_basic/clickhouse/stdout.txt @@ -1,5 +1,7 @@ [ { + "name": "GetValue", + "cmd": ":one", "ast": { "Stmt": { "DistinctClause": null, @@ -35,8 +37,8 @@ "Larg": null, "Rarg": null }, - "StmtLocation": 24, - "StmtLen": 0 + "StmtLocation": 0, + "StmtLen": 32 } } ] diff --git a/internal/engine/clickhouse/analyze_test.go b/internal/engine/clickhouse/analyze_test.go new file mode 100644 index 0000000000..82563223e3 --- /dev/null +++ b/internal/engine/clickhouse/analyze_test.go @@ -0,0 +1,94 @@ +package clickhouse_test + +import ( + "strings" + "testing" + + "github.com/sqlc-dev/sqlc/internal/core" + "github.com/sqlc-dev/sqlc/internal/core/analyzer" + "github.com/sqlc-dev/sqlc/internal/engine/clickhouse" +) + +func analyzeOne(t *testing.T, ddl, query string) core.PrepareResult { + t.Helper() + cat, err := core.New(clickhouse.Dialect()) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { cat.Close() }) + + if err := clickhouse.LoadSchema(cat, ddl); err != nil { + t.Fatalf("load schema: %v", err) + } + + stmts, err := clickhouse.NewParser().Parse(strings.NewReader(query)) + if err != nil { + t.Fatalf("parse query: %v", err) + } + if len(stmts) != 1 { + t.Fatalf("expected 1 stmt, got %d", len(stmts)) + } + res, err := analyzer.Prepare(cat, stmts[0].Raw) + if err != nil { + t.Fatalf("analyze: %v", err) + } + return res +} + +func colByName(res core.PrepareResult, name string) (core.Column, bool) { + for _, c := range res.Columns { + if c.Name == name { + return c, true + } + } + return core.Column{}, false +} + +const eventsDDL = ` +CREATE TABLE events ( + id UInt64, + name String, + tag Nullable(String), + amount Decimal(18, 4) +) ENGINE = MergeTree ORDER BY id +` + +func TestClickHouseSelectColumns(t *testing.T) { + res := analyzeOne(t, eventsDDL, `SELECT id, name, tag FROM events`) + + if len(res.Columns) != 3 { + t.Fatalf("got %d cols, want 3: %+v", len(res.Columns), res.Columns) + } + + id, ok := colByName(res, "id") + if !ok || id.DataType != "uint64" || !id.NotNull { + t.Errorf("id: %+v (ok=%v)", id, ok) + } + name, ok := colByName(res, "name") + if !ok || name.DataType != "string" || !name.NotNull { + t.Errorf("name: %+v (ok=%v)", name, ok) + } + tag, ok := colByName(res, "tag") + if !ok || tag.DataType != "string" || tag.NotNull { + t.Errorf("tag: want string/nullable, got %+v (ok=%v)", tag, ok) + } + + for _, c := range res.Columns { + if c.SourceClassOID == 0 || c.SourceAttributeOID == 0 { + t.Errorf("col %s missing source binding: %+v", c.Name, c) + } + } +} + +func TestClickHouseSelectStar(t *testing.T) { + res := analyzeOne(t, eventsDDL, `SELECT * FROM events`) + if len(res.Columns) != 4 { + t.Fatalf("got %d cols, want 4: %+v", len(res.Columns), res.Columns) + } + want := []string{"id", "name", "tag", "amount"} + for i, w := range want { + if res.Columns[i].Name != w { + t.Errorf("col %d: got %q, want %q", i, res.Columns[i].Name, w) + } + } +} diff --git a/internal/engine/clickhouse/convert.go b/internal/engine/clickhouse/convert.go index ba2817e2bb..bc4053bb10 100644 --- a/internal/engine/clickhouse/convert.go +++ b/internal/engine/clickhouse/convert.go @@ -1,6 +1,7 @@ package clickhouse import ( + "fmt" "strconv" "strings" @@ -485,8 +486,8 @@ func (c *cc) convertFunctionCall(n *chast.FunctionCall) *ast.FuncCall { Funcname: &ast.List{ Items: []ast.Node{&ast.String{Str: n.Name}}, }, - Location: n.Pos().Offset, - AggDistinct: n.Distinct, + Location: n.Pos().Offset, + AggDistinct: n.Distinct, } // Convert arguments @@ -826,14 +827,7 @@ func (c *cc) convertColumnDeclaration(n *chast.ColumnDeclaration) *ast.ColumnDef if n.Type != nil { colDef.TypeName = &ast.TypeName{ - Name: n.Type.Name, - } - // Handle type parameters (e.g., Decimal(10, 2)) - if len(n.Type.Parameters) > 0 { - colDef.TypeName.Typmods = &ast.List{} - for _, param := range n.Type.Parameters { - colDef.TypeName.Typmods.Items = append(colDef.TypeName.Typmods.Items, c.convertExpr(param)) - } + Name: renderDataType(n.Type), } } @@ -855,6 +849,36 @@ func (c *cc) convertColumnDeclaration(n *chast.ColumnDeclaration) *ast.ColumnDef return colDef } +func renderDataType(dt *chast.DataType) string { + if dt == nil { + return "" + } + if len(dt.Parameters) == 0 { + return dt.Name + } + parts := make([]string, 0, len(dt.Parameters)) + for _, p := range dt.Parameters { + parts = append(parts, renderTypeParam(p)) + } + return dt.Name + "(" + strings.Join(parts, ", ") + ")" +} + +func renderTypeParam(e chast.Expression) string { + switch v := e.(type) { + case *chast.DataType: + return renderDataType(v) + case *chast.Literal: + if v.Source != "" { + return v.Source + } + return fmt.Sprintf("%v", v.Value) + case *chast.Identifier: + return strings.Join(v.Parts, ".") + default: + return "" + } +} + func (c *cc) convertUpdateQuery(n *chast.UpdateQuery) *ast.UpdateStmt { rv := &ast.RangeVar{ Relname: &n.Table, diff --git a/internal/engine/clickhouse/parse.go b/internal/engine/clickhouse/parse.go index 282089f31d..fa77fd91ef 100644 --- a/internal/engine/clickhouse/parse.go +++ b/internal/engine/clickhouse/parse.go @@ -30,35 +30,96 @@ func (p *Parser) Parse(r io.Reader) ([]ast.Statement, error) { } var stmts []ast.Statement + loc := 0 for _, stmt := range stmtNodes { + start := stmt.Pos().Offset - 1 + if start < loc { + start = loc + } + end := statementEnd(blob, start) + converter := &cc{} out := converter.convert(stmt) if _, ok := out.(*ast.TODO); ok { + loc = end continue } - // Get position information from the statement - pos := stmt.Pos() - end := stmt.End() - stmtLen := end.Offset - pos.Offset - stmts = append(stmts, ast.Statement{ Raw: &ast.RawStmt{ Stmt: out, - StmtLocation: pos.Offset, - StmtLen: stmtLen, + StmtLocation: loc, + StmtLen: end - loc, }, }) + loc = end } return stmts, nil } +func statementEnd(blob []byte, start int) int { + for i := start; i < len(blob); i++ { + switch blob[i] { + case '\'', '"', '`': + i = skipQuoted(blob, i) + case '-': + if i+1 < len(blob) && blob[i+1] == '-' { + i = skipLineComment(blob, i) + } + case '#': + i = skipLineComment(blob, i) + case '/': + if i+1 < len(blob) && blob[i+1] == '*' { + i = skipBlockComment(blob, i) + } + case ';': + return i + 1 + } + } + return len(blob) +} + +func skipQuoted(blob []byte, i int) int { + q := blob[i] + for j := i + 1; j < len(blob); j++ { + switch blob[j] { + case '\\': + j++ + case q: + if j+1 < len(blob) && blob[j+1] == q { + j++ + continue + } + return j + } + } + return len(blob) - 1 +} + +func skipLineComment(blob []byte, i int) int { + for j := i; j < len(blob); j++ { + if blob[j] == '\n' { + return j + } + } + return len(blob) - 1 +} + +func skipBlockComment(blob []byte, i int) int { + for j := i + 2; j < len(blob); j++ { + if blob[j] == '*' && j+1 < len(blob) && blob[j+1] == '/' { + return j + 1 + } + } + return len(blob) - 1 +} + // https://clickhouse.com/docs/en/sql-reference/syntax#comments func (p *Parser) CommentSyntax() source.CommentSyntax { return source.CommentSyntax{ - Dash: true, // -- comment - SlashStar: true, // /* comment */ - Hash: true, // # comment (ClickHouse supports this) + Dash: true, // -- comment + SlashStar: true, // /* comment */ + Hash: true, // # comment (ClickHouse supports this) } } diff --git a/internal/engine/clickhouse/schema.go b/internal/engine/clickhouse/schema.go new file mode 100644 index 0000000000..cb37a15bdf --- /dev/null +++ b/internal/engine/clickhouse/schema.go @@ -0,0 +1,194 @@ +package clickhouse + +import ( + "fmt" + "strings" + + "github.com/sqlc-dev/sqlc/internal/core" + "github.com/sqlc-dev/sqlc/internal/sql/ast" +) + +func LoadSchema(cat *core.Catalog, ddl string) error { + stmts, err := NewParser().Parse(strings.NewReader(ddl)) + if err != nil { + return fmt.Errorf("clickhouse: parse schema: %w", err) + } + for _, s := range stmts { + if err := Apply(cat, s.Raw); err != nil { + return err + } + } + return nil +} + +func Apply(cat *core.Catalog, n ast.Node) error { + switch v := n.(type) { + case nil: + return nil + case *ast.RawStmt: + return Apply(cat, v.Stmt) + case *ast.List: + for _, it := range v.Items { + if err := Apply(cat, it); err != nil { + return err + } + } + return nil + case *ast.CreateTableStmt: + return applyCreateTable(cat, v) + case *ast.DropTableStmt: + return applyDropTable(cat, v) + } + return nil +} + +func applyCreateTable(cat *core.Catalog, stmt *ast.CreateTableStmt) error { + if stmt.Name == nil { + return fmt.Errorf("clickhouse: create table with nil name") + } + nsOID, err := resolveOrCreateNamespace(cat, stmt.Name.Schema) + if err != nil { + return err + } + if _, err := cat.ClassOID(nsOID, stmt.Name.Name); err == nil { + if stmt.IfNotExists { + return nil + } + return fmt.Errorf("clickhouse: relation %q already exists", stmt.Name.Name) + } + classOID, err := cat.CreateClass(nsOID, stmt.Name.Name, "r") + if err != nil { + return err + } + for i, col := range stmt.Cols { + if col == nil || col.TypeName == nil { + return fmt.Errorf("clickhouse: column %d on %q: missing type", i+1, stmt.Name.Name) + } + typeName, _, _ := unwrapType(col.TypeName) + typeOID, err := resolveOrCreateType(cat, typeName) + if err != nil { + return fmt.Errorf("clickhouse: column %s.%s: %w", stmt.Name.Name, col.Colname, err) + } + if err := cat.CreateAttributeSpec(core.AttributeSpec{ + ClassOID: classOID, + Name: col.Colname, + TypeOID: typeOID, + Num: i + 1, + NotNull: col.IsNotNull || col.PrimaryKey, + IsPrimaryKey: col.PrimaryKey, + DeclType: col.TypeName.Name, + }); err != nil { + return fmt.Errorf("clickhouse: attr %s.%s: %w", stmt.Name.Name, col.Colname, err) + } + } + return nil +} + +func applyDropTable(cat *core.Catalog, stmt *ast.DropTableStmt) error { + for _, tn := range stmt.Tables { + if tn == nil { + continue + } + nsOID, err := cat.NamespaceOID(nsName(tn.Schema)) + if err != nil { + if stmt.IfExists { + continue + } + return err + } + classOID, err := cat.ClassOID(nsOID, tn.Name) + if err != nil { + if stmt.IfExists { + continue + } + return fmt.Errorf("clickhouse: drop table %q: %w", tn.Name, err) + } + if err := cat.DropClass(classOID); err != nil { + return fmt.Errorf("clickhouse: drop table %q: %w", tn.Name, err) + } + } + return nil +} + +func nsName(schema string) string { + if schema == "" { + return "public" + } + return schema +} + +func resolveOrCreateNamespace(cat *core.Catalog, schema string) (int64, error) { + name := nsName(schema) + if oid, err := cat.NamespaceOID(name); err == nil { + return oid, nil + } + return cat.CreateNamespace(name) +} + +func resolveOrCreateType(cat *core.Catalog, name string) (int64, error) { + canonical := strings.ToLower(strings.TrimSpace(name)) + if canonical == "" { + canonical = "nothing" + } + if oid, err := cat.TypeOID(canonical); err == nil { + return oid, nil + } + return cat.CreateType(canonical, 0) +} + +func unwrapType(tn *ast.TypeName) (name string, isArray, nullable bool) { + return unwrapTypeString(tn.Name) +} + +func unwrapTypeString(s string) (name string, isArray, nullable bool) { + base, args := splitType(s) + switch strings.ToLower(base) { + case "nullable": + if len(args) == 1 { + inner, arr, _ := unwrapTypeString(args[0]) + return inner, arr, true + } + return strings.ToLower(base), false, true + case "lowcardinality": + if len(args) == 1 { + return unwrapTypeString(args[0]) + } + return strings.ToLower(base), false, false + case "array": + if len(args) == 1 { + inner, _, nul := unwrapTypeString(args[0]) + return inner, true, nul + } + return strings.ToLower(base), true, false + default: + return strings.ToLower(base), false, false + } +} + +func splitType(s string) (base string, args []string) { + s = strings.TrimSpace(s) + open := strings.IndexByte(s, '(') + if open < 0 || !strings.HasSuffix(s, ")") { + return s, nil + } + base = strings.TrimSpace(s[:open]) + inner := s[open+1 : len(s)-1] + depth, start := 0, 0 + for i := 0; i < len(inner); i++ { + switch inner[i] { + case '(': + depth++ + case ')': + depth-- + case ',': + if depth == 0 { + args = append(args, strings.TrimSpace(inner[start:i])) + start = i + 1 + } + } + } + if last := strings.TrimSpace(inner[start:]); last != "" { + args = append(args, last) + } + return base, args +} diff --git a/internal/engine/clickhouse/seed.go b/internal/engine/clickhouse/seed.go new file mode 100644 index 0000000000..986838285c --- /dev/null +++ b/internal/engine/clickhouse/seed.go @@ -0,0 +1,51 @@ +package clickhouse + +import ( + "fmt" + + "github.com/sqlc-dev/sqlc/internal/core" +) + +func Dialect() core.Option { + return core.WithSeed(Seed) +} + +type chType struct { + name string + category string +} + +var clickhouseTypes = []chType{ + {"UInt8", "N"}, {"UInt16", "N"}, {"UInt32", "N"}, {"UInt64", "N"}, + {"UInt128", "N"}, {"UInt256", "N"}, + {"Int8", "N"}, {"Int16", "N"}, {"Int32", "N"}, {"Int64", "N"}, + {"Int128", "N"}, {"Int256", "N"}, + {"Float32", "N"}, {"Float64", "N"}, {"BFloat16", "N"}, + {"Decimal", "N"}, {"Decimal32", "N"}, {"Decimal64", "N"}, + {"Decimal128", "N"}, {"Decimal256", "N"}, + {"Bool", "B"}, + {"String", "S"}, {"FixedString", "S"}, {"UUID", "S"}, + {"Date", "D"}, {"Date32", "D"}, {"DateTime", "D"}, {"DateTime64", "D"}, + {"IPv4", "S"}, {"IPv6", "S"}, {"JSON", "U"}, + {"Enum8", "U"}, {"Enum16", "U"}, + {"Nullable", "U"}, {"LowCardinality", "U"}, {"Array", "A"}, + {"Map", "U"}, {"Tuple", "U"}, {"Nested", "U"}, {"Nothing", "U"}, +} + +func Seed(cat *core.Catalog) error { + dialectOID, err := cat.CreateDialect("clickhouse") + if err != nil { + return err + } + for _, t := range clickhouseTypes { + if _, err := cat.CreateTypeSpec(core.TypeSpec{ + Name: t.name, + Typtype: "b", + Category: t.category, + DialectOID: dialectOID, + }); err != nil { + return fmt.Errorf("seed clickhouse type %q: %w", t.name, err) + } + } + return nil +}