From 387b9fa12ff5c522921d1346029021af2d67fcff Mon Sep 17 00:00:00 2001 From: "shaoyue.chen" Date: Fri, 18 Sep 2026 20:10:45 +0800 Subject: [PATCH 1/5] Add shared Milvus SQL lifecycle with MySQL and PostgreSQL protocols --- .github/workflows/test.yml | 53 ++++ Readme.md | 212 +++++++++------- config.yaml | 17 +- go.mod | 36 ++- go.sum | 87 ++++++- pkg/bind.go | 196 +++++++++++++++ pkg/commands.go | 199 +++++++++++++++ pkg/config.go | 40 +++ pkg/conn.go | 397 +++-------------------------- pkg/conn_query.go | 132 ++++------ pkg/conn_resultvalue.go | 57 ++--- pkg/ddl.go | 80 ++++-- pkg/engine_test.go | 233 +++++++++++++++++ pkg/indexes.go | 35 +++ pkg/insert.go | 407 +++++++++++++++++++----------- pkg/integration_test.go | 206 +++++++++++++++ pkg/mysql.go | 130 ++++++++++ pkg/parameters.go | 150 +++++++++++ pkg/postgres.go | 501 +++++++++++++++++++++++++++++++++++++ pkg/protocol_test.go | 285 +++++++++++++++++++++ pkg/select.go | 244 ++++++++++++++---- pkg/select_plan.go | 246 ++++++++++++++++++ pkg/select_test.go | 146 +++++++++++ pkg/server.go | 369 ++++++++------------------- pkg/show.go | 8 +- pkg/types.go | 11 +- testdata/embedEtcd.yaml | 5 + 27 files changed, 3390 insertions(+), 1092 deletions(-) create mode 100644 .github/workflows/test.yml create mode 100644 pkg/bind.go create mode 100644 pkg/commands.go create mode 100644 pkg/engine_test.go create mode 100644 pkg/indexes.go create mode 100644 pkg/integration_test.go create mode 100644 pkg/mysql.go create mode 100644 pkg/parameters.go create mode 100644 pkg/postgres.go create mode 100644 pkg/protocol_test.go create mode 100644 pkg/select_plan.go create mode 100644 pkg/select_test.go create mode 100644 testdata/embedEtcd.yaml diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml new file mode 100644 index 0000000..916b5d3 --- /dev/null +++ b/.github/workflows/test.yml @@ -0,0 +1,53 @@ +name: test +on: + push: + pull_request: +permissions: + contents: read +jobs: + unit: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - uses: actions/setup-go@v5 + with: + go-version-file: go.mod + cache: true + - run: go test -race ./... + - run: go vet ./... + milvus: + runs-on: ubuntu-latest + timeout-minutes: 20 + steps: + - uses: actions/checkout@v4 + - uses: actions/setup-go@v5 + with: + go-version-file: go.mod + cache: true + - name: Start isolated Milvus + run: | + docker run -d --name sqlproxy-milvus --security-opt seccomp:unconfined \ + -e ETCD_USE_EMBED=true -e ETCD_DATA_DIR=/var/lib/milvus/etcd \ + -e ETCD_CONFIG_PATH=/milvus/configs/embedEtcd.yaml \ + -e COMMON_STORAGETYPE=local -e DEPLOY_MODE=STANDALONE \ + -v "$PWD/testdata/embedEtcd.yaml:/milvus/configs/embedEtcd.yaml:ro" \ + -p 127.0.0.1:19530:19530 -p 127.0.0.1:9091:9091 \ + milvusdb/milvus:v2.6.2 milvus run standalone + for attempt in $(seq 1 120); do + if curl -fsS http://127.0.0.1:9091/healthz; then exit 0; fi + sleep 2 + done + docker logs --tail 100 sqlproxy-milvus + exit 1 + - name: Test both protocols against Milvus + env: + MILVUS_TEST_ADDR: localhost:19530 + run: go test -race ./pkg -count=1 -v -coverprofile=coverage.out + - name: Milvus logs on failure + if: failure() + run: docker logs --tail 150 sqlproxy-milvus || true + - uses: actions/upload-artifact@v4 + if: always() + with: + name: coverage + path: coverage.out diff --git a/Readme.md b/Readme.md index 1e550dd..b843319 100644 --- a/Readme.md +++ b/Readme.md @@ -1,101 +1,131 @@ -# milvus-sql-proxy -Milvus SQL Proxy is a proxy service that translates SQL queries into Milvus[https://milvus.io] grpc requests. Make integration with milvus easier. +# Milvus SQL Proxy -It can function as a client side sidecar proxy, or a server side proxy, so that you can use various sql driver to connect to milvus. +Connect MySQL or PostgreSQL clients to Milvus using a shared SQL execution layer. +Both protocols can run together; each connection owns its own Milvus client and +selected database. This is a Milvus adapter, not a relational database: joins, +transactions and arbitrary SQL expressions are not supported. -It's still in early alpha stage. +## Run +Requires Go 1.23+ and Milvus (integration-tested with 2.6.2). -## Get Started +```sh +go run ./cmd -config config.yaml +mysql -h 127.0.0.1 -P 3306 -u root default +psql 'postgresql://root@127.0.0.1:5432/default?sslmode=disable' +``` -1. install milvus-lite: `python3 -m pip install milvus` -2. run milvus-server: `milvus-server` -3. run milvus-sql-proxy: `go run cmd/milvus-sql.go` -4. run mysql client: `mysql -u root -h 127.0.0.1 -P 3306` +`mode` is `mysql`, `postgres`, or `both`; omitted mode preserves the original +MySQL-only behavior. `addr` is the MySQL listener and `postgresAddr` is the +PostgreSQL listener. The sample runs both, bound to loopback. Set `user` and +`password` to authenticate SQL clients; these are separate from Milvus credentials. +For network access, configure a password and frontend `tlsCert`/`tlsKey`, and +require TLS in clients. PostgreSQL uses password authentication; without TLS the +password travels in cleartext. The proxy does not expose per-user Milvus RBAC. +`milvus.tlsSecure` enables TLS to Milvus; HTTPS endpoints also enable it. + +Commands time out after `queryTimeoutSeconds` (default 30). Increase this for +large synchronous index builds, load, or flush operations. A timed-out Milvus +mutation may have been accepted; inspect its state before retrying. + +## Supported SQL + +The same SQL subset is available on both ports. MySQL text and prepared/binary +queries are supported. PostgreSQL simple and extended queries (Parse, Bind, +Describe, Execute, Sync), `$n` parameters, text/binary scalar values and quoted +identifiers are supported, including pgx's default prepared-statement mode. +This does not emulate PostgreSQL system catalogs, ORM migrations or all dialect +syntax. MySQL uses `?` parameters; PostgreSQL uses `$1`, `$2`, etc. + +| Area | Commands | +| --- | --- | +| Databases | `SHOW DATABASES`, `CREATE DATABASE db`, `USE db`, `DROP DATABASE db` | +| Collections | `SHOW TABLES`, `CREATE TABLE`, `DESCRIBE table`, `DROP TABLE` | +| Indexes | `CREATE INDEX`, `SHOW INDEXES FROM table`, `DROP INDEX name ON table` | +| Memory/storage | `LOAD TABLE table`, `RELEASE TABLE table`, `FLUSH TABLE table` | +| Partitions | `CREATE/DROP/LOAD/RELEASE PARTITION name ON table`, `SHOW PARTITIONS FROM table` | +| Writes | Multi-row `INSERT`, full-row `UPSERT` (also `REPLACE`), filtered `DELETE` | +| Reads | Field projection, `*`, scalar filters, `count(*)`, pagination, dense vector ANN | + +Field types: `BOOL`/`BOOLEAN` (also `TINYINT(1)`), `TINYINT`, `SMALLINT`, `INT`, +`BIGINT`, `FLOAT`, `DOUBLE`, `VARCHAR(n)`, `JSON`, `VECTOR(dim)` (float32). +Collections need one inline `BIGINT` or `VARCHAR` primary key and at least one +vector field. Auto-ID uses `BIGINT AUTO_INCREMENT PRIMARY KEY`. All non-auto-ID +fields must be supplied; SQL NULL/nullable fields and default values are not +supported. Upsert requires an explicit primary key; auto-ID upserts are rejected. +Vectors are JSON numeric arrays; dimensions and every row are validated before +sending a write. JSON columns accept a string containing valid JSON. -## Supported Commands -- [x] show databases -- [x] create database -- [x] use database -- [x] drop database -- [x] show tables -- [x] create table - - [x] auto increment primary key - - [ ] with index -- [x] drop table -- [x] insert -- [ ] create index -- [ ] load -- [ ] release -- [x] select - - [x] scalar query - - [] vector search -- [ ] delete +```sql +CREATE DATABASE demo; +USE demo; +CREATE TABLE documents ( + id BIGINT PRIMARY KEY, + title VARCHAR(256), + published BOOLEAN, + embedding VECTOR(3) +); +CREATE INDEX embedding_idx ON documents (embedding) + USING HNSW WITH (metric_type='COSINE', M=16, efConstruction=200); +INSERT INTO documents VALUES + (1, 'first', true, json_vector('[1,0,0]')), + (2, 'second', false, json_vector('[0,1,0]')); +LOAD TABLE documents; + +SELECT id, title FROM documents WHERE published=true AND id IN (1,2) LIMIT 10; +SELECT count(*) FROM documents WHERE id>=1; +SELECT id, title, _distance FROM documents + WHERE embedding LIKE json_vector('[1,0,0]') AND published=true LIMIT 3; +UPSERT INTO documents VALUES (2, 'updated', true, json_vector('[0,0,1]')); +DELETE FROM documents WHERE id=2; + +RELEASE TABLE documents; +DROP INDEX embedding_idx ON documents; +DROP TABLE documents; +USE default; +DROP DATABASE demo; +``` -## What we want to implement +Create does **not** implicitly index or load. Index methods: `FLAT`, `HNSW`, +`IVF_FLAT`, `AUTOINDEX`, scalar `INVERTED`. Dense metrics: `L2` (default), `IP`, +`COSINE`. Search uses the field's actual index metric, returns nearest neighbors +in Milvus order and optionally `_distance` (distance/similarity according to the +metric). Only one positive vector predicate is allowed, optionally combined with +scalar predicates using `AND`; vector predicates inside `OR` or `NOT` are rejected. +Writes and searches currently target the default partition. + +Scalar predicates: `=`, `!=`, `<>`, `<`, `<=`, `>`, `>=`, `IN`, `NOT IN`, `LIKE`, +`AND`, `OR`, `NOT`, parentheses. `LIMIT count OFFSET offset` and `LIMIT offset,count` +are supported; default limit 100, maximum limit + offset 16384. `LIMIT 0` returns +metadata and no rows. Count has no pagination. Query/search use strong consistency. +SDK v2 does not return deleted row count, so DELETE reports 0 affected rows even +when matching entities were removed. + +Unsupported clauses return errors: aliases, joins, qualified collection names, +DISTINCT, GROUP BY, HAVING, ORDER BY, transactions, UPDATE, NULL predicates, +IF EXISTS / IF NOT EXISTS, schema alterations, multi-statement requests and +arbitrary functions. Database selection is through startup database or `USE`. +Sparse/binary vectors, hybrid search/reranking, RBAC and bulk import are outside +this SQL subset. PostgreSQL cancellation packets are not supported; server query +timeouts and connection/server shutdown bound execution. + +## Test + +```sh +go test -race ./... +go vet ./... +# Point only at a disposable test Milvus; the test creates/removes its own databases. +MILVUS_TEST_ADDR=localhost:19530 go test -race ./pkg -run TestMilvusIntegration -v +``` -To use MySQL protocol to connect to & operate [Milvus](https://milvus.io/) +CI runs the lifecycle through both `database/sql` MySQL and pgx clients against a +standalone Milvus container, including authentication failures and rejected SQL. +Unit tests use isolated session mocks and real local protocol listeners. -```sql -show databases; -/* output: -+-----------+ -| DATABASES | -+-----------+ -| default | -+-----------+ -1 rows in set (0.01 sec) -*/ - --- create database -create database mydb; -/* output: -Query OK, 1 row affected (0.04 sec) -*/ - -use mydb; -/* output: -Database changed -*/ - --- create collection -create table test ( - id bigint AUTO_INCREMENT PRIMARY KEY, - name varchar(255), - vec vector(32)); -- NOTE: zilliz cloud supports >=32 dimension - --- create vector index -create HNSW("L2") index vec_idx on test (vec); - --- insert vector -insert into test (name, vec) values - ("jack", json_vector("[1.0,2,3,4,5,6,7,8,9,10,11,12,13,14,15,16,17,18,19,20,21,22,23,24,25,26,27,28,29,30,31,32]")), - ("tom", json_vector("[2,3,4,5,6,7,8,9,10,11,12,13,14,15,16,17,18,19,20,21,22,23,24,25,26,27,28,29,30,31,32,33]")), - ("lucy", json_vector("[3,4,5,6,7,8,9,10,11,12,13,14,15,16,17,18,19,20,21,22,23,24,25,26,27,28,29,30,31,32,33,34]")), - ("lily", json_vector("[4,5,6,7,8,9,10,11,12,13,14,15,16,17,18,19,20,21,22,23,24,25,26,27,28,29,30,31,32,33,34,35]")), - ("nova", json_vector("[5,6,7,8,9,10,11,12,13,14,15,16,17,18,19,20,21,22,23,24,25,26,27,28,29,30,31,32,33,34,35,36]")), - ("peter", json_vector("[6,7,8,9,10,11,12,13,14,15,16,17,18,19,20,21,22,23,24,25,26,27,28,29,30,31,32,34,35,36,37,38]")), - ("john", json_vector("[7,8,9,10,11,12,13,14,15,16,17,18,19,20,21,22,23,24,25,26,27,28,29,30,32,33,34,35,36,37,38,39]")), - ("jason", json_vector("[8,9,10,11,12,13,14,15,16,17,18,19,20,21,22,23,24,25,26,27,29,30,31,32,33,34,35,36,37,38,39,40]")); -/* output: -Query OK, 8 row affected (0.10 sec) -*/ - --- ANN search -select id from test where vec like json_vector("[1,2,3,4,5,6,7,8,9,10,11,12,13,14,15,16,17,18,19,20,21,22,23,24,25,26,27,28,29,30,31,32]") limit 3; --- simple query -select * from test where id=1; - --- delete data -delete from test where id = 1; - --- delete collection -truncate table test; - --- drop database -drop database mydb; -``` +## Credits -# Thanks -- http://github.com/xwb1989/sqlparser for the brilliant sql parser. -- http://github.com/flike/kingshard for their sql proxy server framework. +SQL grammar: [haorenfsa/sqlparser](https://github.com/haorenfsa/sqlparser), forked +from [xwb1989/sqlparser](https://github.com/xwb1989/sqlparser). MySQL protocol: +[go-mysql](https://github.com/go-mysql-org/go-mysql). PostgreSQL protocol: +[pgx](https://github.com/jackc/pgx). Original result helpers derive from +[kingshard](https://github.com/flike/kingshard). diff --git a/config.yaml b/config.yaml index 7f57792..cc9f93c 100644 --- a/config.yaml +++ b/config.yaml @@ -1,11 +1,18 @@ logLevel: info -addr: "0.0.0.0:3306" +# mysql, postgres, or both. The legacy addr setting is the MySQL listener. +mode: both +addr: "127.0.0.1:3306" +postgresAddr: "127.0.0.1:5432" +user: root +password: "" +queryTimeoutSeconds: 30 +# Optional frontend TLS (both protocols). +# tlsCert: server.crt +# tlsKey: server.key milvus: addr: "http://localhost:19530" - - # for zilliz cloud: - # addr: https://: + # Zilliz Cloud: use an https endpoint plus apiKey, or user/pass. apiKey: "" - # you can also use user + password user: "" pass: "" + tlsSecure: false diff --git a/go.mod b/go.mod index 6190121..33b6aed 100644 --- a/go.mod +++ b/go.mod @@ -1,36 +1,56 @@ module github.com/haorenfsa/milvus-sql-proxy -go 1.18 +go 1.23.0 + +toolchain go1.24.10 require ( github.com/flike/kingshard v0.0.0-20200829024017-f17b39394746 + github.com/go-mysql-org/go-mysql v1.13.0 + github.com/go-sql-driver/mysql v1.7.1 + github.com/jackc/pgx/v5 v5.7.6 + github.com/milvus-io/milvus-proto/go-api/v2 v2.3.3 github.com/milvus-io/milvus-sdk-go/v2 v2.3.3 github.com/pkg/errors v0.9.1 github.com/xwb1989/sqlparser v0.0.0-20180606152119-120387863bf2 - gopkg.in/yaml.v2 v2.2.5 + gopkg.in/yaml.v2 v2.2.8 ) require ( + filippo.io/edwards25519 v1.1.0 // indirect github.com/cockroachdb/errors v1.9.1 // indirect github.com/cockroachdb/logtags v0.0.0-20211118104740-dabe8e521a4f // indirect github.com/cockroachdb/redact v1.1.3 // indirect github.com/getsentry/sentry-go v0.12.0 // indirect + github.com/goccy/go-json v0.10.2 // indirect github.com/gogo/protobuf v1.3.2 // indirect github.com/golang/protobuf v1.5.2 // indirect + github.com/google/uuid v1.3.0 // indirect github.com/grpc-ecosystem/go-grpc-middleware v1.3.0 // indirect - github.com/kr/pretty v0.3.0 // indirect + github.com/jackc/pgpassfile v1.0.0 // indirect + github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect + github.com/klauspost/compress v1.17.8 // indirect + github.com/kr/pretty v0.3.1 // indirect github.com/kr/text v0.2.0 // indirect - github.com/milvus-io/milvus-proto/go-api/v2 v2.3.3 // indirect - github.com/rogpeppe/go-internal v1.8.1 // indirect + github.com/pingcap/errors v0.11.5-0.20250318082626-8f80e5cb09ec // indirect + github.com/pingcap/log v1.1.1-0.20241212030209-7e3ff8601a2a // indirect + github.com/pingcap/tidb/pkg/parser v0.0.0-20250421232622-526b2c79173d // indirect + github.com/rogpeppe/go-internal v1.10.0 // indirect + github.com/shopspring/decimal v1.2.0 // indirect github.com/tidwall/gjson v1.14.4 // indirect github.com/tidwall/match v1.1.1 // indirect github.com/tidwall/pretty v1.2.0 // indirect - golang.org/x/net v0.17.0 // indirect - golang.org/x/sys v0.13.0 // indirect - golang.org/x/text v0.13.0 // indirect + go.uber.org/atomic v1.11.0 // indirect + go.uber.org/multierr v1.11.0 // indirect + go.uber.org/zap v1.27.0 // indirect + golang.org/x/crypto v0.37.0 // indirect + golang.org/x/net v0.21.0 // indirect + golang.org/x/sys v0.32.0 // indirect + golang.org/x/text v0.24.0 // indirect google.golang.org/genproto v0.0.0-20220503193339-ba3ae3f07e29 // indirect google.golang.org/grpc v1.48.0 // indirect google.golang.org/protobuf v1.30.0 // indirect + gopkg.in/natefinch/lumberjack.v2 v2.2.1 // indirect ) replace github.com/xwb1989/sqlparser => github.com/haorenfsa/sqlparser v0.1.0 diff --git a/go.sum b/go.sum index 2046dd7..c72b2ba 100644 --- a/go.sum +++ b/go.sum @@ -1,5 +1,7 @@ cloud.google.com/go v0.26.0/go.mod h1:aQUYkXzVsufM+DwF1aE+0xfcU+56JwCaLick0ClmMTw= cloud.google.com/go v0.34.0/go.mod h1:aQUYkXzVsufM+DwF1aE+0xfcU+56JwCaLick0ClmMTw= +filippo.io/edwards25519 v1.1.0 h1:FNf4tywRC1HmFuKW5xopWpigGjJKiJSV0Cqo0cJWDaA= +filippo.io/edwards25519 v1.1.0/go.mod h1:BxyFTGdWcka3PhytdK4V28tE5sGfRvvvRV7EaN4VDT4= github.com/AndreasBriese/bbloom v0.0.0-20190306092124-e2d15f34fcf9/go.mod h1:bOvUY6CB00SOBii9/FifXqc0awNKxLFCL/+pkDPuyl8= github.com/BurntSushi/toml v0.3.1/go.mod h1:xHWCNGjB5oqiDr8zfno3MHue2Ht5sIBksp03qcyfWMU= github.com/CloudyKit/fastprinter v0.0.0-20200109182630-33d98a066a53/go.mod h1:+3IMCy2vIlbG1XG/0ggNQv0SvxCAIpPM5b1nCz56Xno= @@ -14,6 +16,7 @@ github.com/alecthomas/units v0.0.0-20190717042225-c3de453c63f4/go.mod h1:ybxpYRF github.com/antihax/optional v1.0.0/go.mod h1:uupD/76wgC+ih3iEmQUL+0Ugr19nfwCT1kdvxnR2qWY= github.com/armon/consul-api v0.0.0-20180202201655-eb2c6b5be1b6/go.mod h1:grANhF5doyWs3UAsr3K4I6qtAmlQcZDesFNEHPZAzj8= github.com/aymerick/raymond v2.0.3-0.20180322193309-b565731e1464+incompatible/go.mod h1:osfaiScAUVup+UC9Nfq76eWqDhXlp+4UYaA8uhTBO6g= +github.com/benbjohnson/clock v1.1.0/go.mod h1:J11/hYXuz8f4ySSvYwY0FKfm+ezbsZBKZxNJlLklBHA= github.com/beorn7/perks v0.0.0-20180321164747-3a771d992973/go.mod h1:Dwedo/Wpr24TaqPxmxbtue+5NUziq4I4S80YR8gNf3Q= github.com/beorn7/perks v1.0.0/go.mod h1:KWe93zE9D1o94FZ5RNwFwVgaQK1VOXiVxmqh+CedLV8= github.com/beorn7/perks v1.0.1/go.mod h1:G2ZrVWU2WbWT9wwq4/hrbKbnv/1ERSJQ0ibhJ6rlkpw= @@ -71,15 +74,22 @@ github.com/go-check/check v0.0.0-20180628173108-788fd7840127/go.mod h1:9ES+weclK github.com/go-errors/errors v1.0.1 h1:LUHzmkK3GUKUrL/1gfBUxAHzcev3apQlezX/+O7ma6w= github.com/go-errors/errors v1.0.1/go.mod h1:f4zRHt4oKfwPJE5k8C9vpYG+aDHdBFUsgrm6/TyX73Q= github.com/go-faker/faker/v4 v4.1.0 h1:ffuWmpDrducIUOO0QSKSF5Q2dxAht+dhsT9FvVHhPEI= +github.com/go-faker/faker/v4 v4.1.0/go.mod h1:uuNc0PSRxF8nMgjGrrrU4Nw5cF30Jc6Kd0/FUTTYbhg= github.com/go-kit/kit v0.8.0/go.mod h1:xBxKIO96dXMWWy0MnWVtmwkA9/13aqxPnvrjFYMA2as= github.com/go-kit/kit v0.9.0/go.mod h1:xBxKIO96dXMWWy0MnWVtmwkA9/13aqxPnvrjFYMA2as= github.com/go-logfmt/logfmt v0.3.0/go.mod h1:Qt1PoO58o5twSAckw1HlFXLmHsOX5/0LbT9GBnD5lWE= github.com/go-logfmt/logfmt v0.4.0/go.mod h1:3RMwSq7FuexP4Kalkev3ejPJsZTpXXBr9+V4qmtdjCk= github.com/go-martini/martini v0.0.0-20170121215854-22fa46961aab/go.mod h1:/P9AEU963A2AYjv4d1V5eVL1CQbEJq6aCNHDDjibzu8= +github.com/go-mysql-org/go-mysql v1.13.0 h1:Hlsa5x1bX/wBFtMbdIOmb6YzyaVNBWnwrb8gSIEPMDc= +github.com/go-mysql-org/go-mysql v1.13.0/go.mod h1:FQxw17uRbFvMZFK+dPtIPufbU46nBdrGaxOw0ac9MFs= +github.com/go-sql-driver/mysql v1.7.1 h1:lUIinVbN1DY0xBg0eMOzmmtGoHwWBbvnWubQUrtU8EI= +github.com/go-sql-driver/mysql v1.7.1/go.mod h1:OXbVy3sEdcQ2Doequ6Z5BW6fXNQTmx+9S1MCJN5yJMI= github.com/go-stack/stack v1.8.0/go.mod h1:v0f6uXyyMGvRgIKkXu+yp6POWl0qKG85gN/melR3HDY= github.com/gobwas/httphead v0.0.0-20180130184737-2c6c146eadee/go.mod h1:L0fX3K22YWvt/FAX9NnzrNzcI4wNYi9Yku4O0LKYflo= github.com/gobwas/pool v0.2.0/go.mod h1:q8bcK0KcYlCgd9e7WYLm9LpyS+YeLd8JVDW6WezmKEw= github.com/gobwas/ws v1.0.2/go.mod h1:szmBTxLgaFppYjEmNtny/v3w89xOydFnnZMcgRRu/EM= +github.com/goccy/go-json v0.10.2 h1:CrxCmQqYDkv1z7lO7Wbh2HN93uovUHgrECaO5ZrCXAU= +github.com/goccy/go-json v0.10.2/go.mod h1:6MelG93GURQebXPDq3khkgXZkazVtN9CRI+MGFi0w8I= github.com/gogo/googleapis v0.0.0-20180223154316-0cd9801be74a/go.mod h1:gf4bu3Q80BeJ6H1S1vYPm8/ELATdvryBaNFGgqEef3s= github.com/gogo/googleapis v1.4.1/go.mod h1:2lpHqI5OcWCtVElxXnPt+s8oJvMpySlOyM6xDCrzib4= github.com/gogo/protobuf v1.1.1/go.mod h1:r8qH/GZQm5c6nD/R0oafs1akxWv10x8SbQlK7atdtwQ= @@ -117,6 +127,8 @@ github.com/google/go-cmp v0.5.6/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/ github.com/google/go-querystring v1.0.0/go.mod h1:odCYkC5MyYFN7vkCjXpyrEuKhc/BUO6wN/zVPAxq5ck= github.com/google/gofuzz v1.0.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/M65Eg= github.com/google/uuid v1.1.2/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +github.com/google/uuid v1.3.0 h1:t6JiXgmwXMjEs8VusXIJk2BXHsn+wx8BZdTaoZ5fu7I= +github.com/google/uuid v1.3.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= github.com/gopherjs/gopherjs v0.0.0-20181017120253-0766667cb4d1/go.mod h1:wJfORRmW1u3UXTncJ5qlYoELFm8eSnnEO6hX4iZ3EWY= github.com/gorilla/websocket v1.4.1/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE= github.com/grpc-ecosystem/go-grpc-middleware v1.3.0 h1:+9834+KizmvFV7pXQGSXQTsaWhq2GjuNUt0aUU0YBYw= @@ -135,6 +147,14 @@ github.com/iris-contrib/go.uuid v2.0.0+incompatible/go.mod h1:iz2lgM/1UnEf1kP0L/ github.com/iris-contrib/jade v1.1.3/go.mod h1:H/geBymxJhShH5kecoiOCSssPX7QWYH7UaeZTSWddIk= github.com/iris-contrib/pongo2 v0.0.1/go.mod h1:Ssh+00+3GAZqSQb30AvBRNxBx7rf0GqwkjqxNd0u65g= github.com/iris-contrib/schema v0.0.1/go.mod h1:urYA3uvUNG1TIIjOSCzHr9/LmbQo8LrOcOqfqxa4hXw= +github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsIM= +github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg= +github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo= +github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM= +github.com/jackc/pgx/v5 v5.7.6 h1:rWQc5FwZSPX58r1OQmkuaNicxdmExaEz5A2DO2hUuTk= +github.com/jackc/pgx/v5 v5.7.6/go.mod h1:aruU7o91Tc2q2cFp5h4uP3f6ztExVpyVv88Xl/8Vl8M= +github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo= +github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4= github.com/json-iterator/go v1.1.6/go.mod h1:+SdeFBvtyEkXs7REEP0seUULqWtbJapLOCVDaaPEHmU= github.com/json-iterator/go v1.1.7/go.mod h1:KdQUCv79m/52Kvf8AW2vK1V8akMuk1QjK/uOdHXbAo4= github.com/json-iterator/go v1.1.9/go.mod h1:KdQUCv79m/52Kvf8AW2vK1V8akMuk1QjK/uOdHXbAo4= @@ -150,12 +170,15 @@ github.com/kisielk/errcheck v1.5.0/go.mod h1:pFxgyoBC7bSaBwPgfKdkLd5X25qrDl4LWUI github.com/kisielk/gotool v1.0.0/go.mod h1:XhKaO+MFFWcvkIS/tQcRk01m1F5IRFswLeQ+oQHNcck= github.com/klauspost/compress v1.8.2/go.mod h1:RyIbtBH6LamlWaDj8nUwkbUhJ87Yi3uG0guNDohfE1A= github.com/klauspost/compress v1.9.7/go.mod h1:RyIbtBH6LamlWaDj8nUwkbUhJ87Yi3uG0guNDohfE1A= +github.com/klauspost/compress v1.17.8 h1:YcnTYrq7MikUT7k0Yb5eceMmALQPYBW/Xltxn0NAMnU= +github.com/klauspost/compress v1.17.8/go.mod h1:Di0epgTjJY877eYKx5yC51cX2A2Vl2ibi7bDH9ttBbw= github.com/klauspost/cpuid v1.2.1/go.mod h1:Pj4uuM528wm8OyEC2QMXAi2YiTZ96dNQPGgoMS4s3ek= github.com/konsorten/go-windows-terminal-sequences v1.0.1/go.mod h1:T0+1ngSBFLxvqU3pZ+m/2kptfBszLMUkC4ZK/EgS/cQ= github.com/kr/logfmt v0.0.0-20140226030751-b84e30acd515/go.mod h1:+0opPa2QZZtGFBFZlji/RkVcI2GknAs/DXo4wKdlNEc= github.com/kr/pretty v0.1.0/go.mod h1:dAy3ld7l9f0ibDNOQOHHMYYIIbhfbHSm3C4ZsoJORNo= -github.com/kr/pretty v0.3.0 h1:WgNl7dwNpEZ6jJ9k1snq4pZsg7DOEN8hP9Xw0Tsjwk0= github.com/kr/pretty v0.3.0/go.mod h1:640gp4NfQd8pI5XOwp5fnNeVWj67G7CFk/SaSQn7NBk= +github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE= +github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk= github.com/kr/pty v1.1.1/go.mod h1:pFQYn66WHrOpPYNljwOMqo10TkYh1fy3cYio2l3bCsQ= github.com/kr/text v0.1.0/go.mod h1:4Jbv+DJW3UT/LiOwJeYQe1efqtUx/iVham/4vfdArNI= github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= @@ -197,8 +220,14 @@ github.com/onsi/ginkgo v1.10.3/go.mod h1:lLunBs/Ym6LB5Z9jYTR76FiuTmxDTDusOGeTQH+ github.com/onsi/gomega v1.7.1/go.mod h1:XdKZgCCFLUoM/7CFJVPcG8C1xQ1AJ0vpAezJrB7JYyY= github.com/opentracing/opentracing-go v1.1.0/go.mod h1:UkNAQd3GIcIGf0SeVgPpRdFStlNbqXla1AfSYxPUl2o= github.com/pelletier/go-toml v1.2.0/go.mod h1:5z9KED0ma1S8pY6P1sdut58dfprrGBbd/94hg7ilaic= -github.com/pingcap/errors v0.11.4 h1:lFuQV/oaUMGcD2tqt+01ROSmJs75VG1ToEOkZIZ4nE4= +github.com/pingcap/errors v0.11.0/go.mod h1:Oi8TUi2kEtXXLMJk9l1cGmz20kV3TaQ0usTwv5KuLY8= github.com/pingcap/errors v0.11.4/go.mod h1:Oi8TUi2kEtXXLMJk9l1cGmz20kV3TaQ0usTwv5KuLY8= +github.com/pingcap/errors v0.11.5-0.20250318082626-8f80e5cb09ec h1:3EiGmeJWoNixU+EwllIn26x6s4njiWRXewdx2zlYa84= +github.com/pingcap/errors v0.11.5-0.20250318082626-8f80e5cb09ec/go.mod h1:X2r9ueLEUZgtx2cIogM0v4Zj5uvvzhuuiu7Pn8HzMPg= +github.com/pingcap/log v1.1.1-0.20241212030209-7e3ff8601a2a h1:WIhmJBlNGmnCWH6TLMdZfNEDaiU8cFpZe3iaqDbQ0M8= +github.com/pingcap/log v1.1.1-0.20241212030209-7e3ff8601a2a/go.mod h1:ORfBOFp1eteu2odzsyaxI+b8TzJwgjwyQcGhI+9SfEA= +github.com/pingcap/tidb/pkg/parser v0.0.0-20250421232622-526b2c79173d h1:3Ej6eTuLZp25p3aH/EXdReRHY12hjZYs3RrGp7iLdag= +github.com/pingcap/tidb/pkg/parser v0.0.0-20250421232622-526b2c79173d/go.mod h1:+8feuexTKcXHZF/dkDfvCwEyBAmgb4paFc3/WeYV2eE= github.com/pkg/diff v0.0.0-20210226163009-20ebb0f2a09e/go.mod h1:pJLUxLENpZxwdsKMEsNbx1VGcRFpLqf3715MtcvvzbA= github.com/pkg/errors v0.8.0/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0= github.com/pkg/errors v0.8.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0= @@ -219,12 +248,16 @@ github.com/prometheus/procfs v0.0.2/go.mod h1:TjEm7ze935MbeOT/UhFTIMYKhuLP4wbCsT github.com/prometheus/procfs v0.0.5/go.mod h1:4A/X28fw3Fc593LaREMrKMqOKvUAntwMDaekg4FpcdQ= github.com/rogpeppe/fastuuid v1.2.0/go.mod h1:jVj6XXZzXRy/MSR5jhDC/2q6DgLz+nrA6LYCDYWNEvQ= github.com/rogpeppe/go-internal v1.6.1/go.mod h1:xXDCJY+GAPziupqXw64V24skbSoqbTEfhy4qGm1nDQc= -github.com/rogpeppe/go-internal v1.8.1 h1:geMPLpDpQOgVyCg5z5GoRwLHepNdb71NXb67XFkP+Eg= github.com/rogpeppe/go-internal v1.8.1/go.mod h1:JeRgkft04UBgHMgCIwADu4Pn6Mtm5d4nPKWu0nJ5d+o= +github.com/rogpeppe/go-internal v1.9.0/go.mod h1:WtVeX8xhTBvf0smdhujwtBcq4Qrzq/fJaraNFVN+nFs= +github.com/rogpeppe/go-internal v1.10.0 h1:TMyTOH3F/DB16zRVcYyreMH6GnZZrwQVAoYjRBZyWFQ= +github.com/rogpeppe/go-internal v1.10.0/go.mod h1:UQnix2H7Ngw/k4C5ijL5+65zddjncjaFoBhdsK/akog= github.com/russross/blackfriday v1.5.2/go.mod h1:JO/DiYxRf+HjHt06OyowR9PTA263kcR/rfWxYHBV53g= github.com/ryanuber/columnize v2.1.0+incompatible/go.mod h1:sm1tb6uqfes/u+d4ooFouqFdy9/2g9QGwK3SQygK0Ts= github.com/schollz/closestmatch v2.1.0+incompatible/go.mod h1:RtP1ddjLong6gTkbtmuhtR2uUrrJOpYzYRvbcPAid+g= github.com/sergi/go-diff v1.0.0/go.mod h1:0CfEIISq7TuYL3j771MWULgwwjU+GofnZX9QAmXWZgo= +github.com/shopspring/decimal v1.2.0 h1:abSATXmQEYyShuxI4/vyW3tV1MrKAJzCZ/0zLUXYbsQ= +github.com/shopspring/decimal v1.2.0/go.mod h1:DKyhrW/HYNuLGql+MJL6WCR6knT2jwCFRcu2hWCYk4o= github.com/shurcooL/sanitized_anchor_name v1.0.0/go.mod h1:1NzhyTcUVG4SuEtjjoZeVRXNmyL/1OwPU0+IJeTBvfc= github.com/sirupsen/logrus v1.2.0/go.mod h1:LxeOpSwHxABJmUn/MG1IvRgCAasNZTLOkJPxbbu5VWo= github.com/sirupsen/logrus v1.4.2/go.mod h1:tLMulIdttU9McNUspp0xgXVQah82FyeX6MwdIuYE2rE= @@ -239,12 +272,14 @@ github.com/spf13/viper v1.3.2/go.mod h1:ZiWeW+zYFKm7srdB9IoDzzZXaJaI5eL9QjNiN/DM github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= github.com/stretchr/objx v0.1.1/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= github.com/stretchr/objx v0.5.0 h1:1zr/of2m5FGMsad5YfcqgdqdWrIhu+EBEJRhR1U7z/c= +github.com/stretchr/objx v0.5.0/go.mod h1:Yh+to48EsGEfYuaHDzXPcE3xhTkx73EhmCGUpEOglKo= github.com/stretchr/testify v1.2.2/go.mod h1:a8OnRcib4nhh0OaRAV+Yts87kKdq0PP7pXfy6kDkUVs= github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI= github.com/stretchr/testify v1.4.0/go.mod h1:j7eGeouHqKxXV5pUuKE4zz7dFj8WfuZ+81PSLYec5m4= github.com/stretchr/testify v1.5.1/go.mod h1:5W2xD1RspED5o8YsWQXVCued0rvSQ+mT+I5cxcmMvtA= github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= -github.com/stretchr/testify v1.8.1 h1:w7B6lhMri9wdJUVmEZPGGhZzrYTPvgJArz7wNPgYKsk= +github.com/stretchr/testify v1.8.4 h1:CcVxjf3Q8PM0mHUKJCdn+eZZtm5yQwehR5yeSVQQcUk= +github.com/stretchr/testify v1.8.4/go.mod h1:sz/lmYIOXD/1dqDmKjjqLyZ2RngseejIcXlSw2iwfAo= github.com/tidwall/gjson v1.14.4 h1:uo0p8EbA09J7RQaflQ1aBRffTR7xedD2bcIVSYxLnkM= github.com/tidwall/gjson v1.14.4/go.mod h1:/wbyibRr2FHMks5tjHJ5F8dMZh3AcwJEMf5vlfC0lxk= github.com/tidwall/match v1.1.1 h1:+Ho715JplO36QYgwN9PGYNhgZvoUSc9X2c80KVTi+GA= @@ -275,8 +310,23 @@ github.com/yuin/goldmark v1.2.1/go.mod h1:3hX8gzYuyVAZsxl0MRgGTJEmQBFcNTphYh9dec github.com/yuin/goldmark v1.3.5/go.mod h1:mwnBkeHKe2W/ZEtQ+71ViKU8L12m81fl3OWwC1Zlc8k= go.opentelemetry.io/proto/otlp v0.7.0/go.mod h1:PqfVotwruBrMGOCsRd/89rSnXhoiJIqeYNgFYFoEGnI= go.uber.org/atomic v1.4.0/go.mod h1:gD2HeocX3+yG+ygLZcrzQJaqmWj9AIm7n08wl/qW/PE= +go.uber.org/atomic v1.6.0/go.mod h1:sABNBOSYdrvTF6hTgEIbc7YasKWGhgEQZyfxyTvoXHQ= +go.uber.org/atomic v1.7.0/go.mod h1:fEN4uk6kAWBTFdckzkM89CLk9XfWZrxpCo0nPH17wJc= +go.uber.org/atomic v1.9.0/go.mod h1:fEN4uk6kAWBTFdckzkM89CLk9XfWZrxpCo0nPH17wJc= +go.uber.org/atomic v1.11.0 h1:ZvwS0R+56ePWxUNi+Atn9dWONBPp/AUETXlHW0DxSjE= +go.uber.org/atomic v1.11.0/go.mod h1:LUxbIzbOniOlMKjJjyPfpl4v+PKK2cNJn91OQbhoJI0= +go.uber.org/goleak v1.1.10/go.mod h1:8a7PlsEVH3e/a/GLqe5IIrQx6GzcnRmZEufDUTk4A7A= +go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto= +go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE= go.uber.org/multierr v1.1.0/go.mod h1:wR5kodmAFQ0UK8QlbwjlSNy0Z68gJhDJUG5sjR94q/0= +go.uber.org/multierr v1.6.0/go.mod h1:cdWPpRnG4AhwMwsgIHip0KRBQjJy5kYEpYjJxpXp9iU= +go.uber.org/multierr v1.7.0/go.mod h1:7EAYxJLBy9rStEaz58O2t4Uvip6FSURkq8/ppBp95ak= +go.uber.org/multierr v1.11.0 h1:blXXJkSxSSfBVBlC76pxqeO+LN3aDfLQo+309xJstO0= +go.uber.org/multierr v1.11.0/go.mod h1:20+QtiLqy0Nd6FdQB9TLXag12DsQkrbs3htMFfDN80Y= go.uber.org/zap v1.10.0/go.mod h1:vwi/ZaCAaUcBkycHslxD9B2zi4UTXhF60s6SWpuDF0Q= +go.uber.org/zap v1.19.0/go.mod h1:xg/QME4nWcxGxrpdeYfq7UvYrLh66cuVKdrbD1XF/NI= +go.uber.org/zap v1.27.0 h1:aJMhYGrd5QSmlpLMr2MftRKl7t8J8PTZPA732ud/XR8= +go.uber.org/zap v1.27.0/go.mod h1:GB2qFLM7cTU87MWRP2mPIjqfIDnGu+VIO4V/SdhGo2E= golang.org/x/crypto v0.0.0-20180904163835-0709b304e793/go.mod h1:6SG95UA2DQfeDnfUPMdvaQW0Q7yPrPDi9nlGo2tz2b4= golang.org/x/crypto v0.0.0-20181203042331-505ab145d0a9/go.mod h1:6SG95UA2DQfeDnfUPMdvaQW0Q7yPrPDi9nlGo2tz2b4= golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w= @@ -286,10 +336,13 @@ golang.org/x/crypto v0.0.0-20191227163750-53104e6ec876/go.mod h1:LzIPMQfyMNhhGPh golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto= golang.org/x/crypto v0.0.0-20210322153248-0c34fe9e7dc2/go.mod h1:T9bdIzuCu7OtxOm1hfPfRQxPLYneinmdGuTeoZ9dtd4= golang.org/x/crypto v0.0.0-20210921155107-089bfa567519/go.mod h1:GvvjBRRGRdwPK5ydBHafDWAxML/pGHZbMvKqRZ5+Abc= +golang.org/x/crypto v0.37.0 h1:kJNSjF/Xp7kU0iB2Z+9viTPMW4EqqsrywMXLJOOsXSE= +golang.org/x/crypto v0.37.0/go.mod h1:vg+k43peMZ0pUMhYmVAWysMK35e6ioLh3wB8ZCAfbVc= golang.org/x/exp v0.0.0-20190121172915-509febef88a4/go.mod h1:CJ0aWSM057203Lf6IL+f9T1iT9GByDxfZKAQTCR3kQA= golang.org/x/lint v0.0.0-20181026193005-c67002cb31c3/go.mod h1:UVdnD1Gm6xHRNCYTkRU2/jEulfH38KcIWyp/GAMgvoE= golang.org/x/lint v0.0.0-20190227174305-5b3e6a55c961/go.mod h1:wehouNa3lNwaWXcvxsM5YxQ5yQlVC4a0KAMCusXpPoU= golang.org/x/lint v0.0.0-20190313153728-d0100b6bd8b3/go.mod h1:6SW0HCj/g11FgYtHlgUYUwCkIfeOF89ocIRzGO/8vkc= +golang.org/x/lint v0.0.0-20190930215403-16217165b5de/go.mod h1:6SW0HCj/g11FgYtHlgUYUwCkIfeOF89ocIRzGO/8vkc= golang.org/x/lint v0.0.0-20210508222113-6edffad5e616/go.mod h1:3xt1FjdF8hUf6vQPIChWIBhFzV8gjjsPE/fR3IyQdNY= golang.org/x/mod v0.1.1-0.20191105210325-c90efee705ee/go.mod h1:QqPTAvyqsEbceGzBzNggFXnrqF1CaUcvgkdR5Ot7KZg= golang.org/x/mod v0.2.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA= @@ -316,8 +369,8 @@ golang.org/x/net v0.0.0-20201021035429-f5854403a974/go.mod h1:sp8m0HH+o8qH0wwXwY golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg= golang.org/x/net v0.0.0-20210405180319-a5a99cb37ef4/go.mod h1:p54w0d4576C0XHj96bSt6lcn1PtDYWL6XObtHCRCNQM= golang.org/x/net v0.0.0-20211008194852-3b03d305991f/go.mod h1:9nx3DQGgdP8bBQD5qxJ1jj9UTztislL4KSBs9R2vV5Y= -golang.org/x/net v0.17.0 h1:pVaXccu2ozPjCXewfr1S7xza/zcXTity9cCdXQYSjIM= -golang.org/x/net v0.17.0/go.mod h1:NxSsAGuq816PNPmqtQdLE42eU2Fs7NoRIZrHJAlaCOE= +golang.org/x/net v0.21.0 h1:AQyQV4dYCvJ7vGmJyKki9+PBdyvhkSd8EIx/qb0AYv4= +golang.org/x/net v0.21.0/go.mod h1:bIjVDfnllIU7BJ2DNgfnXvpSvtn8VRwhlsaeUTyUS44= golang.org/x/oauth2 v0.0.0-20180821212333-d2e6202438be/go.mod h1:N/0e6XlmueqKjAGxoOufVs8QHGRruUQn6yWY3a++T0U= golang.org/x/oauth2 v0.0.0-20200107190931-bf48bf16ab8d/go.mod h1:gOpvHmFTYa4IltrdGE7lF6nIHvwfUNPOp7c8zoXwtLw= golang.org/x/sync v0.0.0-20180314180146-1d60e4601c6f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= @@ -328,6 +381,8 @@ golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJ golang.org/x/sync v0.0.0-20190911185100-cd5d95a43a6e/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20201020160332-67f06af15bc9/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20210220032951-036812b2e83c/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= +golang.org/x/sync v0.13.0 h1:AauUjRAJ9OSnvULf/ARrrVywoJDy0YS2AwQ98I37610= +golang.org/x/sync v0.13.0/go.mod h1:1dzgHSNfp02xaA81J2MS99Qcpr2w7fw1gpm99rleRqA= golang.org/x/sys v0.0.0-20180830151530-49385e6e1522/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= golang.org/x/sys v0.0.0-20180905080454-ebe1bf3edb33/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= golang.org/x/sys v0.0.0-20180909124046-d0be0721c37e/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= @@ -355,8 +410,8 @@ golang.org/x/sys v0.0.0-20210630005230-0f9fa26af87c/go.mod h1:oPkhp1MJrh7nUepCBc golang.org/x/sys v0.0.0-20210927094055-39ccf1dd6fa6/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20211007075335-d3039528d8ac/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20220209214540-3681064d5158/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= -golang.org/x/sys v0.13.0 h1:Af8nKPmuFypiUBjVoU9V20FiaFXOcuZI21p0ycVYYGE= -golang.org/x/sys v0.13.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/sys v0.32.0 h1:s77OFDvIQeibCmezSnk/q6iAfkdiQaJi4VzroCFrN20= +golang.org/x/sys v0.32.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k= golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo= golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ= golang.org/x/text v0.3.2/go.mod h1:bEr9sfX3Q8Zfm5fL9x+3itogRgK3+ptLWKqgva+5dAk= @@ -364,8 +419,8 @@ golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= golang.org/x/text v0.3.5/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= golang.org/x/text v0.3.6/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= golang.org/x/text v0.3.7/go.mod h1:u+2+/6zg+i71rQMx5EYifcz6MCKuco9NR6JIITiCfzQ= -golang.org/x/text v0.13.0 h1:ablQoSUd0tRdKxZewP80B+BaqeKJuVhuRxj/dkrun3k= -golang.org/x/text v0.13.0/go.mod h1:TvPlkZtksWOMsz7fbANvkp4WM8x/WCo/om8BMLbz+aE= +golang.org/x/text v0.24.0 h1:dd5Bzh4yt5KYA8f9CJHCP4FB4D51c2c6JvN37xJJkJ0= +golang.org/x/text v0.24.0/go.mod h1:L8rBsPeo2pSS+xqN0d5u2ikmjtmoJbDBT1b7nHvFCdU= golang.org/x/time v0.0.0-20201208040808-7e3f01d25324/go.mod h1:tRJNPiyCQ0inRvYxbN9jk5I+vvW/OXSQhTDSoE431IQ= golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ= golang.org/x/tools v0.0.0-20181221001348-537d06c36207/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ= @@ -375,6 +430,8 @@ golang.org/x/tools v0.0.0-20190311212946-11955173bddd/go.mod h1:LCzVGOaR6xXOjkQ3 golang.org/x/tools v0.0.0-20190327201419-c70d86f8b7cf/go.mod h1:LCzVGOaR6xXOjkQ3onu1FJEFr0SW1gC7cKk1uF8kGRs= golang.org/x/tools v0.0.0-20190328211700-ab21143f2384/go.mod h1:LCzVGOaR6xXOjkQ3onu1FJEFr0SW1gC7cKk1uF8kGRs= golang.org/x/tools v0.0.0-20190524140312-2c0ae7006135/go.mod h1:RgjU9mgBXZiqYHBnxXauZ1Gv1EHHAz9KjViQ78xBX0Q= +golang.org/x/tools v0.0.0-20191029041327-9cc4af7d6b2c/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo= +golang.org/x/tools v0.0.0-20191108193012-7d206e10da11/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo= golang.org/x/tools v0.0.0-20191115202509-3a792d9c32b2/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo= golang.org/x/tools v0.0.0-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo= golang.org/x/tools v0.0.0-20200130002326-2f3ba24bd6e7/go.mod h1:TB2adYChydJhpapKDTa4BR/hXlZSLoq2Wpct/0txZ28= @@ -410,6 +467,7 @@ google.golang.org/grpc v1.46.0/go.mod h1:vN9eftEi1UMyUsIF80+uQXhHjbXYbm0uXoFCACu google.golang.org/grpc v1.48.0 h1:rQOsyJ/8+ufEDJd/Gdsz7HG220Mh9HAhFHRGnIjda0w= google.golang.org/grpc v1.48.0/go.mod h1:vN9eftEi1UMyUsIF80+uQXhHjbXYbm0uXoFCACuMGWk= google.golang.org/grpc/examples v0.0.0-20220617181431-3e7b97febc7f h1:rqzndB2lIQGivcXdTuY3Y9NBvr70X+y77woofSRluec= +google.golang.org/grpc/examples v0.0.0-20220617181431-3e7b97febc7f/go.mod h1:gxndsbNG1n4TZcHGgsYEfVGnTxqfEdfiDv6/DADXX9o= google.golang.org/protobuf v0.0.0-20200109180630-ec00e32a8dfd/go.mod h1:DFci5gLYBciE7Vtevhsrf46CRTquxDuWsQurQQe4oz8= google.golang.org/protobuf v0.0.0-20200221191635-4d8936d0db64/go.mod h1:kwYJMbMJ01Woi6D6+Kah6886xMZcty6N08ah7+eCXa0= google.golang.org/protobuf v0.0.0-20200228230310-ab0ca4ff8a60/go.mod h1:cfTl7dwQJ+fmap5saPgwCLgHXTUD7jkjRqWcaiX5VyM= @@ -428,24 +486,29 @@ google.golang.org/protobuf v1.30.0/go.mod h1:HV8QOd/L58Z+nl8r43ehVNZIU/HEI6OcFqw gopkg.in/alecthomas/kingpin.v2 v2.2.6/go.mod h1:FMv+mEhP44yOT+4EoQTLFTRgOQ1FBLkstjWtayDeSgw= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= gopkg.in/check.v1 v1.0.0-20180628173108-788fd7840127/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= -gopkg.in/check.v1 v1.0.0-20190902080502-41f04d3bba15 h1:YR8cESwS4TdDjEe65xsg0ogRM/Nc3DYOhEAlW+xobZo= gopkg.in/check.v1 v1.0.0-20190902080502-41f04d3bba15/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk= +gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c/go.mod h1:JHkPIbrfpd72SG/EVd6muEfDQjcINNoR0C8j2r3qZ4Q= gopkg.in/errgo.v2 v2.1.0/go.mod h1:hNsd1EY+bozCKY1Ytp96fpM3vjJbqLJn88ws8XvfDNI= gopkg.in/fsnotify.v1 v1.4.7/go.mod h1:Tz8NjZHkW78fSQdbUxIjBTcgA1z1m8ZHf0WmKUhAMys= gopkg.in/go-playground/assert.v1 v1.2.1/go.mod h1:9RXL0bg/zibRAgZUYszZSwO/z8Y/a8bDuhia5mkpMnE= gopkg.in/go-playground/validator.v8 v8.18.2/go.mod h1:RX2a/7Ha8BgOhfk7j780h4/u/RRjR0eouCJSH80/M2Y= gopkg.in/ini.v1 v1.51.1/go.mod h1:pNLf8WUiyNEtQjuu5G5vTm06TEv9tsIgeAvK8hOrP4k= gopkg.in/mgo.v2 v2.0.0-20180705113604-9856a29383ce/go.mod h1:yeKp02qBN3iKW1OzL3MGk2IdtZzaj7SFntXj72NppTA= +gopkg.in/natefinch/lumberjack.v2 v2.2.1 h1:bBRl1b0OH9s/DuPhuXpNl+VtCaJXFZ5/uEFST95x9zc= +gopkg.in/natefinch/lumberjack.v2 v2.2.1/go.mod h1:YD8tP3GAjkrDg1eZH7EGmyESg/lsYskCTPBJVb9jqSc= gopkg.in/tomb.v1 v1.0.0-20141024135613-dd632973f1e7/go.mod h1:dt/ZhP58zS4L8KSrWDmTeBkI65Dw0HsyUHuEVlX15mw= gopkg.in/yaml.v2 v2.2.1/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= gopkg.in/yaml.v2 v2.2.2/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= gopkg.in/yaml.v2 v2.2.3/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= gopkg.in/yaml.v2 v2.2.4/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= -gopkg.in/yaml.v2 v2.2.5 h1:ymVxjfMaHvXD8RqPRmzHHsB3VvucivSkIAvJFDI5O3c= gopkg.in/yaml.v2 v2.2.5/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= +gopkg.in/yaml.v2 v2.2.8 h1:obN1ZagJSUGI0Ek/LBmuj4SNLPfIny3KsKFopxRdj10= +gopkg.in/yaml.v2 v2.2.8/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= gopkg.in/yaml.v3 v3.0.0-20191120175047-4206685974f2/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= gopkg.in/yaml.v3 v3.0.0-20210107192922-496545a6307b/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= honnef.co/go/tools v0.0.0-20190102054323-c2f93a96b099/go.mod h1:rf3lG4BRIbNafJWhAfAdb/ePZxsR/4RtNHQocxwk9r4= honnef.co/go/tools v0.0.0-20190523083050-ea95bdfd59fc/go.mod h1:rf3lG4BRIbNafJWhAfAdb/ePZxsR/4RtNHQocxwk9r4= diff --git a/pkg/bind.go b/pkg/bind.go new file mode 100644 index 0000000..39872bb --- /dev/null +++ b/pkg/bind.go @@ -0,0 +1,196 @@ +package pkg + +import ( + "fmt" + "strconv" + "strings" +) + +func sqlLiteral(v interface{}) (string, error) { + switch x := v.(type) { + case nil: + return "null", nil + case string: + return "'" + strings.NewReplacer("\\", "\\\\", "'", "\\'", "\x00", "\\0").Replace(x) + "'", nil + case []byte: + return sqlLiteral(string(x)) + case bool: + if x { + return "true", nil + } + return "false", nil + case int: + return strconv.Itoa(x), nil + case int8, int16, int32, int64, uint, uint8, uint16, uint32, uint64, float32, float64: + return fmt.Sprint(x), nil + default: + return "", fmt.Errorf("unsupported parameter type %T", v) + } +} + +// bindSQL scans lexical tokens, never replacing placeholders in literals or +// comments. PostgreSQL quoted identifiers are normalized for the shared parser. +func bindSQL(q, dialect string, args []interface{}, describe bool) (string, int, error) { + var b strings.Builder + count := 0 + for i := 0; i < len(q); { + ch := q[i] + if ch == '\'' || ch == '"' || ch == '`' { + quote := ch + j := i + 1 + var value strings.Builder + closed := false + for j < len(q) { + if q[j] == quote { + if j+1 < len(q) && q[j+1] == quote { + value.WriteByte(quote) + j += 2 + continue + } + j++ + closed = true + break + } + if q[j] == '\\' && dialect == "mysql" && j+1 < len(q) { + j += 2 + continue + } + value.WriteByte(q[j]) + j++ + } + if !closed { + return "", 0, fmt.Errorf("unterminated quoted token") + } + if dialect == "postgres" { + if quote == '"' { + b.WriteByte('`') + b.WriteString(strings.ReplaceAll(value.String(), "`", "``")) + b.WriteByte('`') + } else if quote == '\'' { + lit, _ := sqlLiteral(value.String()) + b.WriteString(lit) + } else { + return "", 0, fmt.Errorf("use double quotes for PostgreSQL identifiers") + } + } else { + b.WriteString(q[i:j]) + } + i = j + continue + } + if i+1 < len(q) && q[i:i+2] == "--" { + j := strings.IndexByte(q[i:], '\n') + if j < 0 { + b.WriteString(q[i:]) + break + } + b.WriteString(q[i : i+j+1]) + i += j + 1 + continue + } + if i+1 < len(q) && q[i:i+2] == "/*" { + j := strings.Index(q[i+2:], "*/") + if j < 0 { + return "", 0, fmt.Errorf("unterminated comment") + } + j += i + 4 + b.WriteString(q[i:j]) + i = j + continue + } + idx := 0 + end := i + 1 + if dialect == "mysql" && ch == '?' { + count++ + idx = count + } + if dialect == "postgres" && ch == '$' { + for end < len(q) && q[end] >= '0' && q[end] <= '9' { + end++ + } + if end == i+1 { + return "", 0, fmt.Errorf("dollar-quoted strings are not supported") + } + n, e := strconv.Atoi(q[i+1 : end]) + if e != nil || n < 1 || n > 65535 { + return "", 0, fmt.Errorf("invalid parameter number") + } + idx = n + if n > count { + count = n + } + } + if idx > 0 { + if describe { + b.WriteString("0") + } else { + if idx > len(args) { + return "", 0, fmt.Errorf("missing parameter %d", idx) + } + v, e := sqlLiteral(args[idx-1]) + if e != nil { + return "", 0, e + } + b.WriteString(v) + } + i = end + continue + } + b.WriteByte(ch) + i++ + } + if !describe && count != len(args) { + return "", 0, fmt.Errorf("expected %d parameters, got %d", count, len(args)) + } + return b.String(), count, nil +} + +// The legacy grammar recognizes BOOL as a keyword but not a column type. +// Rewrite type aliases only outside quoted tokens; VARCHAR contents are intact. +func normalizeTypes(q string) string { + if !strings.HasPrefix(strings.ToUpper(strings.TrimSpace(q)), "CREATE TABLE") { + return q + } + var b strings.Builder + for i := 0; i < len(q); { + if q[i] == '\'' || q[i] == '"' || q[i] == '`' { + quote := q[i] + j := i + 1 + for j < len(q) { + if q[j] == '\\' && j+1 < len(q) { + j += 2 + continue + } + if q[j] == quote { + j++ + if j < len(q) && q[j] == quote { + j++ + continue + } + break + } + j++ + } + b.WriteString(q[i:j]) + i = j + continue + } + if q[i] >= 'a' && q[i] <= 'z' || q[i] >= 'A' && q[i] <= 'Z' || q[i] == '_' { + j := i + 1 + for j < len(q) && (q[j] >= 'a' && q[j] <= 'z' || q[j] >= 'A' && q[j] <= 'Z' || q[j] >= '0' && q[j] <= '9' || q[j] == '_') { + j++ + } + token := q[i:j] + switch strings.ToLower(token) { + case "bool", "boolean": + token = "tinyint(1)" + } + b.WriteString(token) + i = j + continue + } + b.WriteByte(q[i]) + i++ + } + return b.String() +} diff --git a/pkg/commands.go b/pkg/commands.go new file mode 100644 index 0000000..a59d305 --- /dev/null +++ b/pkg/commands.go @@ -0,0 +1,199 @@ +package pkg + +import ( + "encoding/json" + "fmt" + "regexp" + "strings" + + "github.com/milvus-io/milvus-sdk-go/v2/client" + "github.com/milvus-io/milvus-sdk-go/v2/entity" +) + +const ident = "([A-Za-z_][A-Za-z0-9_]*)" + +var lifecycleRE = regexp.MustCompile(`(?i)^(LOAD|RELEASE|FLUSH) (?:TABLE|COLLECTION) ` + ident + `$`) +var describeRE = regexp.MustCompile(`(?i)^(?:DESCRIBE|DESC) ` + ident + `$`) +var createIndexRE = regexp.MustCompile(`(?i)^CREATE INDEX ` + ident + ` ON ` + ident + `\s*\(\s*` + ident + `\s*\) USING (FLAT|IVF_FLAT|HNSW|AUTOINDEX|INVERTED)(?: WITH\s*\((.*)\))?$`) +var dropIndexRE = regexp.MustCompile(`(?i)^DROP INDEX ` + ident + ` ON ` + ident + `$`) +var showIndexesRE = regexp.MustCompile(`(?i)^SHOW INDEXES FROM ` + ident + `$`) +var partitionRE = regexp.MustCompile(`(?i)^(CREATE|DROP|LOAD|RELEASE) PARTITION ` + ident + ` ON ` + ident + `$`) +var showPartitionsRE = regexp.MustCompile(`(?i)^SHOW PARTITIONS FROM ` + ident + `$`) +var upsertRE = regexp.MustCompile(`(?i)^UPSERT\s+INTO\b`) + +type sqlIndex struct { + name, kind string + params map[string]string +} + +func (i sqlIndex) Name() string { return i.name } +func (i sqlIndex) IndexType() entity.IndexType { return entity.IndexType(i.kind) } +func (i sqlIndex) Params() map[string]string { return i.params } + +func (c *ClientConn) rows(names []string, rows [][]interface{}) error { + r, e := c.buildResultset(nil, names, rows) + if e != nil { + return e + } + return c.writeResultset(c.status, r) +} +func (c *ClientConn) handleMilvusCommand(sql string) (bool, error) { + sql = strings.TrimSpace(strings.TrimSuffix(strings.TrimSpace(sql), ";")) + if upsertRE.MatchString(sql) { + if c.describe { + return true, c.writeOK(nil) + } + return true, c.handleQuery(upsertRE.ReplaceAllString(sql, "REPLACE INTO")) + } + if m := lifecycleRE.FindStringSubmatch(sql); m != nil { + if c.describe { + return true, c.writeOK(nil) + } + var e error + switch strings.ToUpper(m[1]) { + case "LOAD": + e = c.upstream.LoadCollection(c.ctx, m[2], false) + case "RELEASE": + e = c.upstream.ReleaseCollection(c.ctx, m[2]) + case "FLUSH": + e = c.upstream.Flush(c.ctx, m[2], false) + } + if e != nil { + return true, e + } + return true, c.writeOK(nil) + } + if m := describeRE.FindStringSubmatch(sql); m != nil { + names := []string{"Field", "Type", "PrimaryKey", "AutoID", "Parameters"} + if c.describe { + r, _ := c.buildResultset(nil, names, nil) + r.Fields[2] = resultField(names[2], entity.FieldTypeBool) + r.Fields[3] = resultField(names[3], entity.FieldTypeBool) + return true, c.writeResultset(c.status, r) + } + schema, e := c.GetCollectinSchema(m[1]) + if e != nil { + return true, e + } + rows := [][]interface{}{} + for _, f := range schema.Fields { + params, _ := json.Marshal(f.TypeParams) + rows = append(rows, []interface{}{f.Name, f.DataType.String(), f.PrimaryKey, f.AutoID, string(params)}) + } + return true, c.rows(names, rows) + } + if m := createIndexRE.FindStringSubmatch(sql); m != nil { + if c.describe { + return true, c.writeOK(nil) + } + params := map[string]string{"index_type": strings.ToUpper(m[4])} + if m[5] != "" { + for _, part := range strings.Split(m[5], ",") { + kv := strings.SplitN(part, "=", 2) + if len(kv) != 2 { + return true, fmt.Errorf("index options must be key=value") + } + key := strings.TrimSpace(kv[0]) + switch key { + case "metric_type", "M", "efConstruction", "nlist": + default: + return true, fmt.Errorf("unsupported index parameter %s", key) + } + if _, ok := params[key]; ok { + return true, fmt.Errorf("duplicate index parameter") + } + params[key] = strings.Trim(strings.TrimSpace(kv[1]), "'\"") + } + } + if params["index_type"] != "INVERTED" { + if params["metric_type"] == "" { + params["metric_type"] = "L2" + } + switch params["metric_type"] { + case "L2", "IP", "COSINE": + default: + return true, fmt.Errorf("metric_type must be L2, IP or COSINE") + } + } + switch params["index_type"] { + case "HNSW": + if params["M"] == "" { + params["M"] = "16" + } + if params["efConstruction"] == "" { + params["efConstruction"] = "200" + } + case "IVF_FLAT": + if params["nlist"] == "" { + params["nlist"] = "128" + } + } + e := c.upstream.CreateIndex(c.ctx, m[2], m[3], sqlIndex{m[1], params["index_type"], params}, false, client.WithIndexName(m[1])) + if e != nil { + return true, e + } + return true, c.writeOK(nil) + } + if m := dropIndexRE.FindStringSubmatch(sql); m != nil { + if c.describe { + return true, c.writeOK(nil) + } + e := c.upstream.DropIndex(c.ctx, m[2], "", client.WithIndexName(m[1])) + if e != nil { + return true, e + } + return true, c.writeOK(nil) + } + if m := showIndexesRE.FindStringSubmatch(sql); m != nil { + names := []string{"Name", "Type", "Parameters"} + if c.describe { + return true, c.rows(names, nil) + } + indexes, e := c.listIndexes(m[1]) + if e != nil { + return true, e + } + rows := [][]interface{}{} + for _, idx := range indexes { + params, _ := json.Marshal(idx.Params()) + rows = append(rows, []interface{}{idx.Name(), string(idx.IndexType()), string(params)}) + } + return true, c.rows(names, rows) + } + if m := partitionRE.FindStringSubmatch(sql); m != nil { + if c.describe { + return true, c.writeOK(nil) + } + var e error + switch strings.ToUpper(m[1]) { + case "CREATE": + e = c.upstream.CreatePartition(c.ctx, m[3], m[2]) + case "DROP": + e = c.upstream.DropPartition(c.ctx, m[3], m[2]) + case "LOAD": + e = c.upstream.LoadPartitions(c.ctx, m[3], []string{m[2]}, false) + case "RELEASE": + e = c.upstream.ReleasePartitions(c.ctx, m[3], []string{m[2]}) + } + if e != nil { + return true, e + } + return true, c.writeOK(nil) + } + if m := showPartitionsRE.FindStringSubmatch(sql); m != nil { + names := []string{"Partition"} + if c.describe { + return true, c.rows(names, nil) + } + ps, e := c.upstream.ShowPartitions(c.ctx, m[1]) + if e != nil { + return true, e + } + rows := [][]interface{}{} + for _, p := range ps { + rows = append(rows, []interface{}{p.Name}) + } + return true, c.rows(names, rows) + } + return false, nil +} diff --git a/pkg/config.go b/pkg/config.go index f6acb33..e2e8abe 100644 --- a/pkg/config.go +++ b/pkg/config.go @@ -1,12 +1,21 @@ package pkg import ( + "fmt" "io/ioutil" "gopkg.in/yaml.v2" ) type Config struct { + Mode string `yaml:"mode"` + PostgresAddr string `yaml:"postgresAddr"` + User string `yaml:"user"` + Password string `yaml:"password"` + TLSCert string `yaml:"tlsCert"` + TLSKey string `yaml:"tlsKey"` + QueryTimeoutSeconds int `yaml:"queryTimeoutSeconds"` + Addr string `yaml:"addr"` Milvus MilvusConfig `yaml:"milvus"` LogLevel string `yaml:"logLevel"` @@ -26,6 +35,9 @@ func ParseConfigData(data []byte) (*Config, error) { if err := yaml.Unmarshal(data, &cfg); err != nil { return nil, err } + if err := cfg.defaults(); err != nil { + return nil, err + } return &cfg, nil } @@ -37,3 +49,31 @@ func ParseConfigFile(fileName string) (*Config, error) { return ParseConfigData(data) } + +func (c *Config) defaults() error { + if c.Mode == "" { + c.Mode = "mysql" + } + if c.Mode != "mysql" && c.Mode != "postgres" && c.Mode != "both" { + return fmt.Errorf("mode must be mysql, postgres or both") + } + if c.Addr == "" { + c.Addr = "127.0.0.1:3306" + } + if c.PostgresAddr == "" { + c.PostgresAddr = "127.0.0.1:5432" + } + if c.User == "" { + c.User = "root" + } + if c.QueryTimeoutSeconds == 0 { + c.QueryTimeoutSeconds = 30 + } + if c.QueryTimeoutSeconds < 0 { + return fmt.Errorf("queryTimeoutSeconds must be positive") + } + if (c.TLSCert == "") != (c.TLSKey == "") { + return fmt.Errorf("tlsCert and tlsKey must both be provided") + } + return nil +} diff --git a/pkg/conn.go b/pkg/conn.go index a23543f..6850c8e 100644 --- a/pkg/conn.go +++ b/pkg/conn.go @@ -1,393 +1,58 @@ -// partially copied & changed from : https://github.com/flike/kingshard/blob/master/proxy/server/conn.go -// Copyright 2016 The kingshard Authors. All rights reserved. -// -// Licensed under the Apache License, Version 2.0 (the "License"): you may -// not use this file except in compliance with the License. You may obtain -// a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, software -// distributed under the License is distributed on an "AS IS" BASIS, WITHOUT -// WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. See the -// License for the specific language governing permissions and limitations -// under the License. - package pkg import ( - "bytes" "context" - "encoding/binary" - "fmt" - "math/rand" - "net" - "runtime" - "sync" - - "github.com/flike/kingshard/core/golog" - "github.com/flike/kingshard/core/hack" "github.com/flike/kingshard/mysql" "github.com/milvus-io/milvus-sdk-go/v2/client" - "github.com/milvus-io/milvus-sdk-go/v2/entity" + "sync" ) -// client <-> Milvus +// ClientConn is a SQL session. Each session owns its Milvus client because +// UsingDatabase changes client state. Transport adapters never share sessions. type ClientConn struct { - ctx context.Context - sync.Mutex - - pkg *mysql.PacketIO - - c net.Conn - - upstream client.Client - tableSchema map[string]entity.Schema - // proxy *Server - - capability uint32 - + ctx context.Context + upstream client.Client + db string connectionId uint32 - - status uint16 - collation mysql.CollationId - charset string - - user string - db string - - salt []byte - - closed bool - - lastInsertId int64 + status uint16 affectedRows int64 - - stmtId uint32 - - stmts map[uint32]*Stmt //prepare相关,client端到proxy的stmt - - configVer uint32 //check config version for reload online -} - -var DEFAULT_CAPABILITY uint32 = mysql.CLIENT_LONG_PASSWORD | mysql.CLIENT_LONG_FLAG | - mysql.CLIENT_CONNECT_WITH_DB | mysql.CLIENT_PROTOCOL_41 | - mysql.CLIENT_TRANSACTIONS | mysql.CLIENT_SECURE_CONNECTION - -var baseConnId uint32 = rand.Uint32() - -func (c *ClientConn) Handshake() error { - if err := c.writeInitialHandshake(); err != nil { - golog.Error("server", "Handshake", err.Error(), - c.connectionId, "msg", "send initial handshake error") - return err - } - - if err := c.readHandshakeResponse(); err != nil { - golog.Error("server", "readHandshakeResponse", - err.Error(), c.connectionId, - "msg", "read Handshake Response error") - return err - } - - if err := c.writeOK(nil); err != nil { - golog.Error("server", "readHandshakeResponse", - "write ok fail", - c.connectionId, "error", err.Error()) - return err - } - - c.pkg.Sequence = 0 - return nil -} - -func (c *ClientConn) Close() error { - if c.closed { - return nil - } - - c.upstream.Close() - c.c.Close() - - c.closed = true - - return nil -} - -func (c *ClientConn) writeInitialHandshake() error { - data := make([]byte, 4, 128) - - //min version 10 - data = append(data, 10) - - //server version[00] - data = append(data, mysql.ServerVersion...) - data = append(data, 0) - - //connection id - data = append(data, byte(c.connectionId), byte(c.connectionId>>8), byte(c.connectionId>>16), byte(c.connectionId>>24)) - - //auth-plugin-data-part-1 - data = append(data, c.salt[0:8]...) - - //filter [00] - data = append(data, 0) - - //capability flag lower 2 bytes, using default capability here - data = append(data, byte(DEFAULT_CAPABILITY), byte(DEFAULT_CAPABILITY>>8)) - - //charset, utf-8 default - data = append(data, uint8(mysql.DEFAULT_COLLATION_ID)) - - //status - data = append(data, byte(c.status), byte(c.status>>8)) - - //below 13 byte may not be used - //capability flag upper 2 bytes, using default capability here - data = append(data, byte(DEFAULT_CAPABILITY>>16), byte(DEFAULT_CAPABILITY>>24)) - - //filter [0x15], for wireshark dump, value is 0x15 - data = append(data, 0x15) - - //reserved 10 [00] - data = append(data, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0) - - //auth-plugin-data-part-2 - data = append(data, c.salt[8:]...) - - //filter [00] - data = append(data, 0) - - return c.writePacket(data) -} - -func (c *ClientConn) readPacket() ([]byte, error) { - return c.pkg.ReadPacket() + result *mysql.Result + describe bool } -func (c *ClientConn) writePacket(data []byte) error { - return c.pkg.WritePacket(data) +func NewSession(ctx context.Context, upstream client.Client) *ClientConn { + return &ClientConn{ctx: ctx, upstream: upstream, db: "default", status: mysql.SERVER_STATUS_AUTOCOMMIT} } -func (c *ClientConn) writePacketBatch(total, data []byte, direct bool) ([]byte, error) { - return c.pkg.WritePacketBatch(total, data, direct) -} - -func (c *ClientConn) readHandshakeResponse() error { - data, err := c.readPacket() - - if err != nil { - return err +func (c *ClientConn) Execute(ctx context.Context, sql string) (*mysql.Result, error) { + c.ctx = ctx + c.result = nil + if err := c.handleQuery(sql); err != nil { + return nil, err } - - pos := 0 - - //capability - c.capability = binary.LittleEndian.Uint32(data[:4]) - pos += 4 - - //skip max packet size - pos += 4 - - //charset, skip, if you want to use another charset, use set names - //c.collation = CollationId(data[pos]) - pos++ - - //skip reserved 23[00] - pos += 23 - - //user name - c.user = string(data[pos : pos+bytes.IndexByte(data[pos:], 0)]) - - pos += len(c.user) + 1 - - //auth length and auth - authLen := int(data[pos]) - pos++ - // auth := data[pos : pos+authLen] - - //check user - // TODO: - // if _, ok := c.proxy.users[c.user]; !ok { - // golog.Error("ClientConn", "readHandshakeResponse", "error", 0, - // "auth", auth, - // "client_user", c.user, - // "config_set_user", c.user, - // "password", c.proxy.users[c.user]) - // return mysql.NewDefaultError(mysql.ER_ACCESS_DENIED_ERROR, c.user, c.c.RemoteAddr().String(), "Yes") - // } - - //check password - // TODO: - // checkAuth := mysql.CalcPassword(c.salt, []byte(c.proxy.users[c.user])) - // if !bytes.Equal(auth, checkAuth) { - // golog.Error("ClientConn", "readHandshakeResponse", "error", 0, - // "auth", auth, - // "checkAuth", checkAuth, - // "client_user", c.user, - // "config_set_user", c.user, - // "password", c.proxy.users[c.user]) - // return mysql.NewDefaultError(mysql.ER_ACCESS_DENIED_ERROR, c.user, c.c.RemoteAddr().String(), "Yes") - // } - - pos += authLen - - var db string - if c.capability&mysql.CLIENT_CONNECT_WITH_DB > 0 { - if len(data[pos:]) == 0 { - return nil - } - - db = string(data[pos : pos+bytes.IndexByte(data[pos:], 0)]) - pos += len(c.db) + 1 - - } - c.db = db - - return nil -} - -func (c *ClientConn) Run() { - defer func() { - r := recover() - if err, ok := r.(error); ok { - const size = 4096 - buf := make([]byte, size) - buf = buf[:runtime.Stack(buf, false)] - - golog.Error("ClientConn", "Run", - err.Error(), 0, - "stack", string(buf)) - } - - c.Close() - }() - for { - data, err := c.readPacket() - - if err != nil { - return - } - - if err := c.dispatch(data); err != nil { - // c.proxy.counter.IncrErrLogTotal() - golog.Error("ClientConn", "Run", - err.Error(), c.connectionId, - ) - c.writeError(err) - if err == mysql.ErrBadConn { - c.Close() - } - } - - if c.closed { - return - } - - c.pkg.Sequence = 0 + if c.result == nil { + c.result = &mysql.Result{} } + return c.result, nil } -func (c *ClientConn) dispatch(data []byte) error { - // c.proxy.counter.IncrClientQPS() - cmd := data[0] - data = data[1:] - - switch cmd { - case mysql.COM_QUIT: - // TODO: - // c.handleRollback() - c.Close() - return nil - case mysql.COM_QUERY, mysql.COM_FIELD_LIST: - return c.handleQuery(hack.String(data)) - case mysql.COM_PING: - return c.writeOK(nil) - case mysql.COM_INIT_DB: - return c.handleUseDB(hack.String(data), nil) - // case mysql.COM_FIELD_LIST: - // return c.handleFieldList(data) - // case mysql.COM_STMT_PREPARE: - // return c.handleStmtPrepare(hack.String(data)) - // case mysql.COM_STMT_EXECUTE: - // return c.handleStmtExecute(data) - // case mysql.COM_STMT_CLOSE: - // return c.handleStmtClose(data) - // case mysql.COM_STMT_SEND_LONG_DATA: - // return c.handleStmtSendLongData(data) - // case mysql.COM_STMT_RESET: - // return c.handleStmtReset(data) - case mysql.COM_SET_OPTION: - return c.writeEOF(0) - default: - msg := fmt.Sprintf("command %d not supported now", cmd) - golog.Error("ClientConn", "dispatch", msg, 0) - return mysql.NewError(mysql.ER_UNKNOWN_ERROR, msg) - } +// Describe obtains SELECT metadata without executing queries or mutations. +func (c *ClientConn) Describe(ctx context.Context, sql string) (*mysql.Result, error) { + c.describe = true + defer func() { c.describe = false }() + return c.Execute(ctx, sql) } - func (c *ClientConn) writeOK(r *mysql.Result) error { if r == nil { - r = &mysql.Result{Status: c.status} - } - data := make([]byte, 4, 32) - - data = append(data, mysql.OK_HEADER) - - data = append(data, mysql.PutLengthEncodedInt(r.AffectedRows)...) - data = append(data, mysql.PutLengthEncodedInt(r.InsertId)...) - - if c.capability&mysql.CLIENT_PROTOCOL_41 > 0 { - data = append(data, byte(r.Status), byte(r.Status>>8)) - data = append(data, 0, 0) + r = &mysql.Result{} } - - return c.writePacket(data) -} - -func (c *ClientConn) writeError(e error) error { - var m *mysql.SqlError - var ok bool - if m, ok = e.(*mysql.SqlError); !ok { - m = mysql.NewError(mysql.ER_UNKNOWN_ERROR, e.Error()) - } - - data := make([]byte, 4, 16+len(m.Message)) - - data = append(data, mysql.ERR_HEADER) - data = append(data, byte(m.Code), byte(m.Code>>8)) - - if c.capability&mysql.CLIENT_PROTOCOL_41 > 0 { - data = append(data, '#') - data = append(data, m.State...) - } - - data = append(data, m.Message...) - - return c.writePacket(data) -} - -func (c *ClientConn) writeEOF(status uint16) error { - data := make([]byte, 4, 9) - - data = append(data, mysql.EOF_HEADER) - if c.capability&mysql.CLIENT_PROTOCOL_41 > 0 { - data = append(data, 0, 0) - data = append(data, byte(status), byte(status>>8)) - } - - return c.writePacket(data) + c.result = r + return nil } - -func (c *ClientConn) writeEOFBatch(total []byte, status uint16, direct bool) ([]byte, error) { - data := make([]byte, 4, 9) - - data = append(data, mysql.EOF_HEADER) - if c.capability&mysql.CLIENT_PROTOCOL_41 > 0 { - data = append(data, 0, 0) - data = append(data, byte(status), byte(status>>8)) +func (c *ClientConn) writeError(err error) error { return err } +func (c *ClientConn) Close() { + if c.upstream != nil { + c.upstream.Close() } - - return c.writePacketBatch(total, data, direct) } diff --git a/pkg/conn_query.go b/pkg/conn_query.go index 0935024..0b6770e 100644 --- a/pkg/conn_query.go +++ b/pkg/conn_query.go @@ -1,61 +1,34 @@ -// partially copied & changed from : https://github.com/flike/kingshard/blob/master/proxy/server/ - -// Copyright 2016 The kingshard Authors. All rights reserved. -// -// Licensed under the Apache License, Version 2.0 (the "License"): you may -// not use this file except in compliance with the License. You may obtain -// a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, software -// distributed under the License is distributed on an "AS IS" BASIS, WITHOUT -// WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. See the -// License for the specific language governing permissions and limitations -// under the License. - package pkg import ( "fmt" - "runtime" - "strings" - - "github.com/flike/kingshard/core/errors" - "github.com/flike/kingshard/core/golog" "github.com/xwb1989/sqlparser" + "regexp" + "strings" ) -/*处理query语句*/ -func (c *ClientConn) handleQuery(sql string) (err error) { - defer func() { - if e := recover(); e != nil { - golog.OutputSql("Error", "err:%v,sql:%s", e, sql) - - if err, ok := e.(error); ok { - const size = 4096 - buf := make([]byte, size) - buf = buf[:runtime.Stack(buf, false)] - - golog.Error("ClientConn", "handleQuery", - err.Error(), 0, - "stack", string(buf), "sql", sql) - } - - err = errors.ErrInternalServer - return - } - }() - - sql = strings.TrimRight(sql, ";") //删除sql语句最后的分号 - golog.Debug("conn", "handleQuery", sql, c.connectionId) - var stmt sqlparser.Statement - stmt, err = sqlparser.Parse(sql) //解析sql语句,得到的stmt是一个interface +func (c *ClientConn) handleQuery(sql string) error { + sql = strings.TrimSpace(sql) + if sql == "" { + return fmt.Errorf("empty SQL statement") + } + if handled, err := c.handleMilvusCommand(sql); handled { + return err + } + if err := validateStatementKind(sql); err != nil { + return err + } + stmt, err := sqlparser.ParseStrictDDL(normalizeTypes(sql)) if err != nil { - golog.Error("conn", "parse", err.Error(), c.connectionId, "sql", sql) return err } - + if c.describe { + switch stmt.(type) { + case *sqlparser.Select, *sqlparser.Show: + default: + return c.writeOK(nil) + } + } switch v := stmt.(type) { case *sqlparser.Show: return c.handleShow(v, nil) @@ -67,40 +40,39 @@ func (c *ClientConn) handleQuery(sql string) (err error) { return c.handleSelect(v, nil) case *sqlparser.Insert: return c.handleInsert(v, nil) - // TODO: - // case *sqlparser.Update: - // return c.handleExec(stmt, nil) - // case *sqlparser.Delete: - // return c.handleExec(stmt, nil) - // case *sqlparser.Set: - // return c.handleSet(v, sql) - // case *sqlparser.Begin: - // return c.handleBegin() - // case *sqlparser.Commit: - // return c.handleCommit() - // case *sqlparser.Rollback: - // return c.handleRollback() - // case *sqlparser.Admin: - // if c.user == "root" { - // return c.handleAdmin(v) - // } - // return fmt.Errorf("statement %T not support now", stmt) - // case *sqlparser.AdminHelp: - // if c.user == "root" { - // return c.handleAdminHelp(v) - // } - // return fmt.Errorf("statement %T not support now", stmt) - // case *sqlparser.UseDB: - // return c.handleUseDB(v.DB) - // case *sqlparser.SimpleSelect: - // return c.handleSimpleSelect(v) - // case *sqlparser.Truncate: - // return c.handleExec(stmt, nil) + case *sqlparser.Delete: + return c.handleDelete(v) + case *sqlparser.Use: + return c.handleUseDB(v.DBName.String(), nil) default: - return fmt.Errorf("statement %T not support now", stmt) + return fmt.Errorf("statement %T is not supported", stmt) } } -func (c *ClientConn) handleExec(stmt sqlparser.Statement, args []interface{}) error { - panic("not implemented") +var createTableSQL = regexp.MustCompile("(?is)^CREATE\\s+TABLE\\s+(?:[A-Za-z_][A-Za-z0-9_]*|`[A-Za-z_][A-Za-z0-9_]*`)\\s*\\(.*\\)$") +var nameOnlySQL = regexp.MustCompile("(?i)^(?:DROP TABLE|CREATE DATABASE|DROP DATABASE|USE)\\s+(?:[A-Za-z_][A-Za-z0-9_]*|`[A-Za-z_][A-Za-z0-9_]*`)$") + +func validateStatementKind(q string) error { + q = strings.TrimSpace(strings.TrimSuffix(strings.TrimSpace(q), ";")) + words := strings.Fields(q) + if len(words) == 0 { + return fmt.Errorf("empty statement") + } + switch strings.ToUpper(words[0]) { + case "CREATE": + if createTableSQL.MatchString(q) || nameOnlySQL.MatchString(q) { + return nil + } + case "DROP", "USE": + if nameOnlySQL.MatchString(q) { + return nil + } + case "SHOW": + if strings.EqualFold(q, "show tables") || strings.EqualFold(q, "show databases") { + return nil + } + case "SELECT", "INSERT", "REPLACE", "DELETE": + return nil + } + return fmt.Errorf("unsupported SQL statement or modifiers") } diff --git a/pkg/conn_resultvalue.go b/pkg/conn_resultvalue.go index e594fca..4c82853 100644 --- a/pkg/conn_resultvalue.go +++ b/pkg/conn_resultvalue.go @@ -22,7 +22,6 @@ import ( "github.com/flike/kingshard/core/errors" "github.com/flike/kingshard/core/hack" "github.com/flike/kingshard/mysql" - "github.com/flike/kingshard/sqlparser" ) func formatValue(value interface{}) ([]byte, error) { @@ -30,6 +29,8 @@ func formatValue(value interface{}) ([]byte, error) { return hack.Slice("NULL"), nil } switch v := value.(type) { + case bool: + return strconv.AppendBool(nil, v), nil case int8: return strconv.AppendInt(nil, int64(v), 10), nil case int16: @@ -67,6 +68,12 @@ func formatValue(value interface{}) ([]byte, error) { func formatField(field *mysql.Field, value interface{}) error { switch value.(type) { + case nil: + field.Type = mysql.MYSQL_TYPE_NULL + return nil + case bool: + field.Charset = 63 + field.Type = mysql.MYSQL_TYPE_TINY case int8, int16, int32, int64, int: field.Charset = 63 field.Type = mysql.MYSQL_TYPE_LONGLONG @@ -126,14 +133,18 @@ func (c *ClientConn) buildResultset(fields []*mysql.Field, names []string, value } else { field := &mysql.Field{} r.Fields[j] = field - r.FieldNames[string(r.Fields[j].Name)] = j field.Name = hack.Slice(names[j]) + r.FieldNames[names[j]] = j if err = formatField(field, value); err != nil { return nil, err } } } + if value == nil { + row = append(row, 0xfb) + continue + } b, err = formatValue(value) if err != nil { return nil, err @@ -151,48 +162,10 @@ func (c *ClientConn) buildResultset(fields []*mysql.Field, names []string, value } func (c *ClientConn) writeResultset(status uint16, r *mysql.Resultset) error { - c.affectedRows = int64(-1) - total := make([]byte, 0, 4096) - data := make([]byte, 4, 512) - var err error - - columnLen := mysql.PutLengthEncodedInt(uint64(len(r.Fields))) - - data = append(data, columnLen...) - total, err = c.writePacketBatch(total, data, false) - if err != nil { - return err - } - - for _, v := range r.Fields { - data = data[0:4] - data = append(data, v.Dump()...) - total, err = c.writePacketBatch(total, data, false) - if err != nil { - return err - } - } - - total, err = c.writeEOFBatch(total, status, false) - if err != nil { - return err - } - - for _, v := range r.RowDatas { - data = data[0:4] - data = append(data, v...) - total, err = c.writePacketBatch(total, data, false) - if err != nil { - return err - } - } - - _, err = c.writeEOFBatch(total, status, true) - return err + c.result = &mysql.Result{Resultset: r, Status: status} + return nil } -var nstring = sqlparser.String - func newEmptyResultset(fields []string) *mysql.Resultset { r := new(mysql.Resultset) r.Fields = make([]*mysql.Field, len(fields)) diff --git a/pkg/ddl.go b/pkg/ddl.go index 7704421..3fc5c0a 100644 --- a/pkg/ddl.go +++ b/pkg/ddl.go @@ -2,6 +2,8 @@ package pkg import ( "fmt" + "strconv" + "strings" "github.com/flike/kingshard/core/golog" "github.com/flike/kingshard/mysql" @@ -23,11 +25,6 @@ func (c *ClientConn) handleDDL(stmt *sqlparser.DDL, args []interface{}) error { } } -type MilvusSchema struct { - *entity.Schema - ShardNum int32 -} - func DDLToMilvusSchema(stmt *sqlparser.DDL) (*MilvusSchema, error) { ret := new(MilvusSchema) ret.Schema = new(entity.Schema) @@ -37,7 +34,10 @@ func DDLToMilvusSchema(stmt *sqlparser.DDL) (*MilvusSchema, error) { if stmt.TableSpec == nil { return nil, errors.Errorf("table spec is nil") } - schema.Description = stmt.TableSpec.Options + if !stmt.NewName.Qualifier.IsEmpty() || stmt.TableSpec.Options != "" || len(stmt.TableSpec.Indexes) > 0 { + return nil, errors.New("qualified tables, table options and table-level constraints are not supported; use inline PRIMARY KEY and CREATE INDEX") + } + schema.Description = "" schema.Fields = make([]*entity.Field, 0, len(stmt.TableSpec.Columns)) for _, col := range stmt.TableSpec.Columns { field, err := columnToMilvusField(col) @@ -46,6 +46,34 @@ func DDLToMilvusSchema(stmt *sqlparser.DDL) (*MilvusSchema, error) { } schema.Fields = append(schema.Fields, field) } + seen := map[string]bool{} + pk := 0 + vectors := 0 + for _, f := range schema.Fields { + if seen[f.Name] { + return nil, errors.New("duplicate field") + } + seen[f.Name] = true + if !fieldIdentifier.MatchString(f.Name) { + return nil, errors.New("invalid field name") + } + if f.PrimaryKey { + pk++ + if f.DataType != entity.FieldTypeInt64 && f.DataType != entity.FieldTypeVarChar { + return nil, errors.New("primary key must be BIGINT or VARCHAR") + } + if f.AutoID && f.DataType != entity.FieldTypeInt64 { + return nil, errors.New("auto-ID requires BIGINT") + } + } + if f.DataType == entity.FieldTypeFloatVector { + vectors++ + } + } + if pk != 1 || vectors == 0 { + return nil, errors.New("collection requires exactly one primary key and at least one vector field") + } + ret.ShardNum = 1 // TODO: shard num // default: ret.ShardNum = 2 return ret, nil @@ -54,13 +82,25 @@ func DDLToMilvusSchema(stmt *sqlparser.DDL) (*MilvusSchema, error) { func columnToMilvusField(col *sqlparser.ColumnDefinition) (*entity.Field, error) { field := new(entity.Field) field.Name = col.Name.String() + if col.Type.Default != nil || col.Type.OnUpdate != nil || col.Type.Unsigned || col.Type.Zerofill || col.Type.Scale != nil || col.Type.Charset != "" || col.Type.Collate != "" { + return nil, errors.New("unsupported column options") + } + var supportType bool - if col.Type.Type == sqlparser.KeywordString(sqlparser.VECTOR) { + if strings.EqualFold(col.Type.Type, "bool") || strings.EqualFold(col.Type.Type, "boolean") || (col.Type.Type == "tinyint" && col.Type.Length != nil && string(col.Type.Length.Val) == "1") { + field.DataType = entity.FieldTypeBool + } else if strings.EqualFold(col.Type.Type, "json") { + field.DataType = entity.FieldTypeJSON + } else if col.Type.Type == sqlparser.KeywordString(sqlparser.VECTOR) { field.DataType = entity.FieldTypeFloatVector if col.Type.Length == nil { return nil, errors.Errorf("vector dim is nil") } golog.Debug("ddl", "columnToMilvusField", "dim", 0, string(col.Type.Length.Val)) + dim, err := strconv.Atoi(string(col.Type.Length.Val)) + if err != nil || dim < 1 || dim > 32768 { + return nil, errors.New("vector dimension must be 1..32768") + } field.TypeParams = map[string]string{ "dim": string(col.Type.Length.Val), } @@ -73,6 +113,10 @@ func columnToMilvusField(col *sqlparser.ColumnDefinition) (*entity.Field, error) if col.Type.Length == nil { return nil, errors.Errorf("varchar max_length must be specified") } + max, err := strconv.Atoi(string(col.Type.Length.Val)) + if err != nil || max < 1 || max > 65535 { + return nil, errors.New("varchar length must be 1..65535") + } field.TypeParams = map[string]string{ "max_length": string(col.Type.Length.Val), } @@ -128,31 +172,19 @@ func (c *ClientConn) handleCreateTable(stmt *sqlparser.DDL, args []interface{}) } golog.Info("ddl", "handleCreateTable", "CreateCollection", 0) // TODO: consistency level - err = c.upstream.CreateCollection(c.ctx, milvusSchema.Schema, milvusSchema.ShardNum, client.WithConsistencyLevel(entity.ClEventually)) - if err != nil { - return mysql.NewError(mysql.ER_CANT_CREATE_TABLE, err.Error()) - } - // TODO: use real - golog.Info("ddl", "handleCreateTable", "CreateIndexIvfFlat", 0) - index, err := entity.NewIndexIvfFlat(entity.IP, 128) - if err != nil { - return mysql.NewError(mysql.ER_CANT_CREATE_TABLE, err.Error()) - } - golog.Info("ddl", "handleCreateTable", "CreateIndex", 0) - err = c.upstream.CreateIndex(c.ctx, milvusSchema.Schema.CollectionName, "vec", index, false) - if err != nil { - return mysql.NewError(mysql.ER_CANT_CREATE_TABLE, err.Error()) - } - golog.Info("ddl", "handleCreateTable", "LoadCollection", 0) - err = c.upstream.LoadCollection(c.ctx, milvusSchema.Schema.CollectionName, false) + err = c.upstream.CreateCollection(c.ctx, milvusSchema.Schema, milvusSchema.ShardNum, client.WithConsistencyLevel(entity.ClStrong)) if err != nil { return mysql.NewError(mysql.ER_CANT_CREATE_TABLE, err.Error()) } + return c.writeOK(nil) } func (c *ClientConn) handleDropTable(stmt *sqlparser.DDL, args []interface{}) error { golog.Info("ddl", "handleDropTable", "DropCollection", 0, stmt.Table.Name.String()) + if !stmt.Table.Qualifier.IsEmpty() || stmt.IfExists { + return fmt.Errorf("qualified DROP and IF EXISTS are unsupported") + } err := c.upstream.DropCollection(c.ctx, stmt.Table.Name.String()) if err != nil { return mysql.NewError(mysql.ER_CANT_DROP_FIELD_OR_KEY, err.Error()) diff --git a/pkg/engine_test.go b/pkg/engine_test.go new file mode 100644 index 0000000..3206047 --- /dev/null +++ b/pkg/engine_test.go @@ -0,0 +1,233 @@ +package pkg + +import ( + "context" + "errors" + "reflect" + "strings" + "testing" + + "github.com/milvus-io/milvus-sdk-go/v2/client" + "github.com/milvus-io/milvus-sdk-go/v2/entity" + "github.com/xwb1989/sqlparser" +) + +func testSchema() *entity.Schema { + return &entity.Schema{CollectionName: "items", Fields: []*entity.Field{{Name: "id", DataType: entity.FieldTypeInt64, PrimaryKey: true}, {Name: "name", DataType: entity.FieldTypeVarChar, TypeParams: map[string]string{"max_length": "16"}}, {Name: "enabled", DataType: entity.FieldTypeBool}, {Name: "score", DataType: entity.FieldTypeDouble}, {Name: "meta", DataType: entity.FieldTypeJSON}, {Name: "vec", DataType: entity.FieldTypeFloatVector, TypeParams: map[string]string{"dim": "3"}}}} +} +func insertAST(t *testing.T, q string) *sqlparser.Insert { + t.Helper() + s, e := sqlparser.Parse(q) + if e != nil { + t.Fatal(e) + } + return s.(*sqlparser.Insert) +} +func TestInsertValidation(t *testing.T) { + c := NewSession(context.Background(), nil) + schema := testSchema() + columns, e := c.FillInsertColumns(schema, insertAST(t, `insert into items values (-5,'one',true,-1.5,'{"x":1}',json_vector('[1,2,3]')), (6,'two',false,2.5,'{}',json_vector('[3,2,1]'))`)) + if e != nil { + t.Fatal(e) + } + if len(columns) != 6 || columns[0].Len() != 2 { + t.Fatal(columns) + } + v, _ := columns[0].Get(0) + if v != int64(-5) { + t.Fatal(v) + } + v, _ = columns[2].Get(0) + if v != true { + t.Fatal(v) + } + for _, q := range []string{ + `insert into items values (1,'x')`, + `insert into items values (1,'x',true,1,'{}',json_vector('[1,2]'))`, + `insert into items values (1,'x',true,1,'{}',json_vector('[null,2,3]'))`, + `insert into items values (1,'x',true,1,'{}',json_vector('["1",2,3]'))`, + `insert into items values (1,'x',true,1,'not json',json_vector('[1,2,3]'))`, + `insert into items values (1,'x',true,abs(1),'{}',json_vector('[1,2,3]'))`, + `insert into items values (1,'x',null,1,'{}',json_vector('[1,2,3]'))`, + `insert into items values (9223372036854775808,'x',true,1,'{}',json_vector('[1,2,3]'))`, + `insert into items values (1,'this string is too long',true,1,'{}',json_vector('[1,2,3]'))`, + `insert into items values (1,42,true,1,'{}',json_vector('[1,2,3]'))`, + `insert into items(id,id) values (1,2)`, `insert into items(nope) values (1)`, `insert into items(id) values (1)`, `insert into items select * from x`, + `insert into items values (1,'x',true,1,'{}',unknown('[1,2,3]'))`, + `insert into items values (1,'x',true,1,'{}',json_vector('[1,2,3]','x'))`, + } { + t.Run(q, func(t *testing.T) { + if _, e := c.FillInsertColumns(schema, insertAST(t, q)); e == nil { + t.Fatal("expected validation error") + } + }) + } + schema.Fields[0].AutoID = true + if _, e := c.FillInsertColumns(schema, insertAST(t, `insert into items(name,enabled,score,meta,vec) values ('x',true,1,'{}',json_vector('[1,2,3]'))`)); e != nil { + t.Fatal(e) + } + if _, e := c.FillInsertColumns(schema, insertAST(t, `replace into items(name,enabled,score,meta,vec) values ('x',true,1,'{}',json_vector('[1,2,3]'))`)); e == nil { + t.Fatal("autoID upsert accepted") + } +} +func TestFieldTypes(t *testing.T) { + for _, tc := range []struct { + kind entity.FieldType + good, bad string + }{ + {entity.FieldTypeInt8, "-128", "128"}, {entity.FieldTypeInt16, "32767", "32768"}, {entity.FieldTypeInt32, "2147483647", "2147483648"}, {entity.FieldTypeInt64, "-9223372036854775808", "9223372036854775808"}, {entity.FieldTypeFloat, "1.25", "'NaN'"}, {entity.FieldTypeDouble, "2.5", "'Inf'"}, {entity.FieldTypeBool, "0", "'invalid'"}, + } { + f := &entity.Field{Name: "f", DataType: tc.kind} + col, e := emptyColumn(f) + if e != nil { + t.Fatal(e) + } + for i, literal := range []string{tc.good, tc.bad} { + stmt := insertAST(t, "insert into x values ("+literal+")") + v, e := fieldValue(f, stmt.Rows.(sqlparser.Values)[0][0]) + if i == 0 { + if e != nil { + t.Fatal(e) + } + if e = col.AppendValue(v); e != nil { + t.Fatal(e) + } + } else if e == nil { + t.Fatalf("accepted %s as %v", literal, tc.kind) + } + } + } +} +func TestDDLValidation(t *testing.T) { + for _, q := range []string{`create table t (id bigint primary key, v vector(3))`, `create table t (id bigint auto_increment primary key, name varchar(16), flag boolean, meta json, v vector(3))`} { + s, e := sqlparser.ParseStrictDDL(normalizeTypes(q)) + if e != nil { + t.Fatal(e) + } + schema, e := DDLToMilvusSchema(s.(*sqlparser.DDL)) + if e != nil { + t.Fatal(e) + } + if schema.CollectionName != "t" { + t.Fatal(schema) + } + } + for _, q := range []string{`create table t (id int primary key, v vector(3))`, `create table t (id bigint, v vector(3))`, `create table t (id bigint primary key, v vector(0))`, `create table t (id bigint primary key, v vector(32769))`, `create table t (id bigint primary key, name varchar(0), v vector(3))`, `create table t (id bigint primary key, id bigint, v vector(3))`, `create table t (id bigint primary key, v vector(3), x bigint default 1)`, `create table t (id bigint primary key)`, `create table t (id bigint primary key, v vector(3), primary key (id))`, `create table t (id varchar(16) auto_increment primary key, v vector(3))`} { + s, e := sqlparser.ParseStrictDDL(normalizeTypes(q)) + if e != nil { + continue + } + if _, e := DDLToMilvusSchema(s.(*sqlparser.DDL)); e == nil { + t.Fatalf("accepted %s", q) + } + } +} +func TestDangerousSQLRejectedBeforeRPC(t *testing.T) { + c := NewSession(context.Background(), nil) + for _, q := range []string{"DROP VIEW items", "DROP TABLE items CASCADE", "DROP TABLE items; DROP TABLE other", "DROP DATABASE IF EXISTS db", "CREATE TABLE IF NOT EXISTS items (id bigint primary key,v vector(3))", "SHOW TABLES FROM other LIKE 'x%'", "SHOW FULL TABLES", "BEGIN", "UPDATE items SET id=2"} { + if _, e := c.Execute(context.Background(), q); e == nil { + t.Fatalf("accepted %s", q) + } + } +} + +type operationMock struct { + client.Client + schema *entity.Schema + calls []string + columns []entity.Column + fail error + db string + queryFilter string +} + +func (m *operationMock) Close() error { return nil } +func (m *operationMock) UsingDatabase(_ context.Context, db string) error { m.db = db; return m.fail } +func (m *operationMock) DescribeCollection(context.Context, string) (*entity.Collection, error) { + if m.fail != nil { + return nil, m.fail + } + return &entity.Collection{Schema: m.schema}, nil +} +func (m *operationMock) CreateCollection(_ context.Context, s *entity.Schema, _ int32, _ ...client.CreateCollectionOption) error { + m.calls = append(m.calls, "create") + m.schema = s + return m.fail +} +func (m *operationMock) DropCollection(context.Context, string, ...client.DropCollectionOption) error { + m.calls = append(m.calls, "drop") + return m.fail +} +func (m *operationMock) Insert(_ context.Context, _, _ string, cols ...entity.Column) (entity.Column, error) { + m.calls = append(m.calls, "insert") + m.columns = cols + return entity.NewColumnInt64("id", []int64{7}), m.fail +} +func (m *operationMock) Upsert(_ context.Context, _, _ string, cols ...entity.Column) (entity.Column, error) { + m.calls = append(m.calls, "upsert") + m.columns = cols + return entity.NewColumnInt64("id", []int64{7}), m.fail +} +func (m *operationMock) Delete(_ context.Context, _, _, filter string) error { + m.calls = append(m.calls, "delete") + m.queryFilter = filter + return m.fail +} +func (m *operationMock) Query(_ context.Context, _ string, _ []string, filter string, fields []string, _ ...client.SearchQueryOptionFunc) (client.ResultSet, error) { + m.queryFilter = filter + result := client.ResultSet{} + for _, n := range fields { + switch n { + case "id", "count(*)": + result = append(result, entity.NewColumnInt64(n, []int64{7})) + case "name": + result = append(result, entity.NewColumnVarChar(n, []string{"one"})) + case "enabled": + result = append(result, entity.NewColumnBool(n, []bool{true})) + case "score": + result = append(result, entity.NewColumnDouble(n, []float64{1.5})) + case "meta": + result = append(result, entity.NewColumnJSONBytes(n, [][]byte{[]byte(`{"x":1}`)})) + case "vec": + result = append(result, entity.NewColumnFloatVector(n, 3, [][]float32{{1, 2, 3}})) + } + } + return result, m.fail +} +func TestMutations(t *testing.T) { + m := &operationMock{schema: testSchema()} + c := NewSession(context.Background(), m) + for _, q := range []string{`create table items (id bigint primary key,v vector(3))`, `drop table items`} { + if _, e := c.Execute(context.Background(), q); e != nil { + t.Fatal(e) + } + } + if !reflect.DeepEqual(m.calls, []string{"create", "drop"}) { + t.Fatal(m.calls) + } + m.schema = testSchema() + for _, verb := range []string{"INSERT", "UPSERT", "REPLACE"} { + r, e := c.Execute(context.Background(), verb+` INTO items VALUES (7,'one',true,1.5,'{}',json_vector('[1,2,3]'))`) + if e != nil { + t.Fatal(e) + } + if r.AffectedRows != 1 { + t.Fatal(r) + } + } + if _, e := c.Execute(context.Background(), "delete from items where id in (1,2)"); e != nil { + t.Fatal(e) + } + if m.queryFilter != "id in [1,2]" { + t.Fatal(m.queryFilter) + } + for _, q := range []string{"delete from items", "delete from items limit 1", "delete from db.items where id=1", "insert ignore into items values (1)", "insert into db.items values (1)"} { + if _, e := c.Execute(context.Background(), q); e == nil { + t.Fatal(q) + } + } + m.fail = errors.New("upstream failed") + if _, e := c.Execute(context.Background(), "drop table items"); e == nil || !strings.Contains(e.Error(), "upstream") { + t.Fatal(e) + } +} diff --git a/pkg/indexes.go b/pkg/indexes.go new file mode 100644 index 0000000..f1cc217 --- /dev/null +++ b/pkg/indexes.go @@ -0,0 +1,35 @@ +package pkg + +import ( + "fmt" + "github.com/milvus-io/milvus-proto/go-api/v2/commonpb" + "github.com/milvus-io/milvus-proto/go-api/v2/milvuspb" + "github.com/milvus-io/milvus-sdk-go/v2/client" + "github.com/milvus-io/milvus-sdk-go/v2/entity" +) + +// SDK v2.3 DescribeIndex validates a nonempty field name, preventing an all-index +// query. Its exported gRPC service uses the same database/auth interceptors. +func (c *ClientConn) listIndexes(table string) ([]entity.Index, error) { + upstream, ok := c.upstream.(*client.GrpcClient) + if !ok { + return nil, fmt.Errorf("index listing requires the Milvus gRPC client") + } + resp, err := upstream.Service.DescribeIndex(c.ctx, &milvuspb.DescribeIndexRequest{CollectionName: table}) + if err != nil { + return nil, err + } + status := resp.GetStatus() + if status.GetErrorCode() == commonpb.ErrorCode_IndexNotExist { + return nil, nil + } + if status == nil || status.GetErrorCode() != commonpb.ErrorCode_Success || status.GetCode() != 0 { + return nil, fmt.Errorf("describe indexes: %s", status.GetReason()) + } + indexes := []entity.Index{} + for _, d := range resp.GetIndexDescriptions() { + params := entity.KvPairsMap(d.Params) + indexes = append(indexes, entity.NewGenericIndex(d.IndexName, entity.IndexType(params["index_type"]), params)) + } + return indexes, nil +} diff --git a/pkg/insert.go b/pkg/insert.go index b225040..9ef0add 100644 --- a/pkg/insert.go +++ b/pkg/insert.go @@ -2,171 +2,292 @@ package pkg import ( "encoding/json" - "reflect" + "fmt" + "math" "strconv" + "strings" - "github.com/cockroachdb/errors" - "github.com/flike/kingshard/core/golog" "github.com/flike/kingshard/mysql" "github.com/milvus-io/milvus-sdk-go/v2/entity" "github.com/xwb1989/sqlparser" ) -func (c *ClientConn) GetCollectinSchema(collectionName string) (*entity.Schema, error) { - // TODO: cache - collection, err := c.upstream.DescribeCollection(c.ctx, collectionName) +func (c *ClientConn) GetCollectinSchema(name string) (*entity.Schema, error) { + col, err := c.upstream.DescribeCollection(c.ctx, name) if err != nil { - return nil, errors.Wrapf(err, "describe collection[%s] failed", collectionName) + return nil, err } - return collection.Schema, nil + if col == nil || col.Schema == nil { + return nil, fmt.Errorf("collection has no schema") + } + return col.Schema, nil } - -func (c *ClientConn) handleInsert(stmt *sqlparser.Insert, args []interface{}) error { - golog.Debug("conn", "handleInsert", "GetCollectinSchema", c.connectionId) - schema, err := c.GetCollectinSchema(stmt.Table.Name.String()) +func (c *ClientConn) handleInsert(s *sqlparser.Insert, _ []interface{}) error { + if !s.Table.Qualifier.IsEmpty() || s.Ignore != "" || len(s.OnDup) > 0 || len(s.Partitions) > 0 { + return fmt.Errorf("qualified tables, IGNORE, ON DUPLICATE KEY and partition clauses are not supported") + } + schema, err := c.GetCollectinSchema(s.Table.Name.String()) if err != nil { - return c.writeError(err) + return err } - golog.Debug("conn", "handleInsert", "FillInsertColumns", c.connectionId) - columns, err := c.FillInsertColumns(schema, stmt) + columns, err := c.FillInsertColumns(schema, s) if err != nil { - return c.writeError(err) + return err + } + var ids entity.Column + if s.Action == sqlparser.ReplaceStr { + ids, err = c.upstream.Upsert(c.ctx, s.Table.Name.String(), "", columns...) + } else { + ids, err = c.upstream.Insert(c.ctx, s.Table.Name.String(), "", columns...) } - - golog.Debug("conn", "handleInsert", "upstream.Insert", c.connectionId) - _, err = c.upstream.Insert(c.ctx, - stmt.Table.Name.String(), - "", // TODO: partition name - columns...) if err != nil { - return c.writeError(err) + return err } - golog.Debug("conn", "handleInsert", "writeOK", c.connectionId, "affectedRows", c.affectedRows) - return c.writeOK(&mysql.Result{ - AffectedRows: uint64(len(stmt.Rows.(sqlparser.Values))), - }) + r := &mysql.Result{AffectedRows: uint64(len(s.Rows.(sqlparser.Values)))} + if ids != nil && ids.Len() > 0 { + if v, e := ids.Get(0); e == nil { + if id, ok := v.(int64); ok && id >= 0 { + r.InsertId = uint64(id) + } + } + } + return c.writeOK(r) } - -func (c *ClientConn) FillInsertColumns(schema *entity.Schema, stmt *sqlparser.Insert) ([]entity.Column, error) { - rowsValues := stmt.Rows.(sqlparser.Values) - if len(rowsValues) == 0 { - return []entity.Column{}, nil +func (c *ClientConn) FillInsertColumns(schema *entity.Schema, s *sqlparser.Insert) ([]entity.Column, error) { + rows, ok := s.Rows.(sqlparser.Values) + if !ok || len(rows) == 0 { + return nil, fmt.Errorf("INSERT requires VALUES rows") } - - var columnSchmaMap = make(map[string]*entity.Field) - for _, field := range schema.Fields { - columnSchmaMap[field.Name] = field + fields := map[string]*entity.Field{} + for _, f := range schema.Fields { + fields[f.Name] = f } - - var columnIndexMap = make(map[string]int) - var columns []entity.Column = make([]entity.Column, 0, len(stmt.Columns)) - if len(stmt.Columns) > 0 { - for columnIdx, columnStmt := range stmt.Columns { - columnName := columnStmt.String() - columnSchema, ok := columnSchmaMap[columnName] - if !ok { - return nil, errors.Errorf("column[%s] not exist", columnName) - + names := []string{} + if len(s.Columns) == 0 { + for _, f := range schema.Fields { + if !f.AutoID { + names = append(names, f.Name) } - golog.Debug("conn", "handleInsert", "insert", c.connectionId, "columnStmt", columnStmt) - columnIndexMap[columnName] = columnIdx - - switch columnSchema.DataType { - case entity.FieldTypeInt32: - var values = make([]int32, len(rowsValues)) - for rowIdx, row := range rowsValues { - columnValues := row[columnIdx] - golog.Debug("conn", "handleInsert", "columnValues.(type)", c.connectionId, "expr", columnValues) - switch expr := columnValues.(type) { - case *sqlparser.SQLVal: - if expr.Type != sqlparser.IntVal { - return nil, errors.Errorf("column[%s] row[%d] type[%s] not [%s]", columnName, rowIdx, expr.Type, "IntVal") - } - val, err := strconv.ParseInt(string(expr.Val), 10, 32) - if err != nil { - return nil, errors.Wrapf(err, "column[%s] row[%d] type[%s] value[%v]", columnName, rowIdx, expr.Type, expr.Val) - } - values[rowIdx] = int32(val) - case *sqlparser.FuncExpr: - return nil, errors.Errorf("column[%s] row[%d] type[%s] not supported", columnName, rowIdx, expr.Name) - } - } - columns = append(columns, entity.NewColumnInt32(columnName, values)) - case entity.FieldTypeInt64: - var values = make([]int64, len(rowsValues)) - for rowIdx, row := range rowsValues { - columnValues := row[columnIdx] - golog.Debug("conn", "handleInsert", "columnValues.(type)", c.connectionId, "expr", columnValues) - switch expr := columnValues.(type) { - case *sqlparser.SQLVal: - if expr.Type != sqlparser.IntVal { - return nil, errors.Errorf("column[%s] row[%d] type[%s] not [%s]", columnName, rowIdx, expr.Type, "IntVal") - } - golog.Debug("conn", "handleInsert", "ParseInt", c.connectionId, "val", expr.Val) - val, err := strconv.ParseInt(string(expr.Val), 10, 64) - if err != nil { - return nil, errors.Wrapf(err, "column[%s] row[%d] type[%s] value[%v]", columnName, rowIdx, expr.Type, expr.Val) - } - values[rowIdx] = val - case *sqlparser.FuncExpr: - return nil, errors.Errorf("column[%s] row[%d] type[%s] not supported", columnName, rowIdx, expr.Name) - } - } - columns = append(columns, entity.NewColumnInt64(columnName, values)) - case entity.FieldTypeVarChar: - var values = make([]string, len(rowsValues)) - for rowIdx, row := range rowsValues { - columnValues := row[columnIdx] - golog.Debug("conn", "handleInsert", "columnValues.(type)", c.connectionId, "expr", columnValues) - switch expr := columnValues.(type) { - case *sqlparser.SQLVal: - if expr.Type != sqlparser.StrVal { - return nil, errors.Errorf("column[%s] row[%d] type[%s] not [%s]", columnName, rowIdx, expr.Type, "StrVal") - } - values[rowIdx] = string(expr.Val) - case *sqlparser.FuncExpr: - return nil, errors.Errorf("column[%s] row[%d] type[%s] not supported", columnName, rowIdx, expr.Name) - } - } - columns = append(columns, entity.NewColumnVarChar(columnName, values)) - case entity.FieldTypeFloatVector: - var values = make([][]float32, len(rowsValues)) - for rowIdx, row := range rowsValues { - columnValues := row[columnIdx] - // golog.Debug("conn", "handleInsert", "columnValues.(type)", c.connectionId, "type", reflect.TypeOf(columnValues)) - switch expr := columnValues.(type) { - case *sqlparser.SQLVal: - return nil, errors.Errorf("column[%s] row[%d] type[%s] not supported", columnName, rowIdx, expr.Type) - case *sqlparser.FuncExpr: - const JSONVectorFuncName = "json_vector" - if expr.Name.String() != JSONVectorFuncName { - return nil, errors.Errorf("column[%s] row[%d] type[%s] not supported", columnName, rowIdx, expr.Name) - } - if len(expr.Exprs) != 1 { - return nil, errors.Errorf("column[%s] row[%d] type[%s] len(exprs) != 1", columnName, rowIdx, expr.Name) - } - vectorExpr, ok := expr.Exprs[0].(*sqlparser.AliasedExpr) - if !ok { - return nil, errors.Errorf("column[%s] row[%d] type[%s] exprs[0] type != *sqlparser.AliasedExpr", columnName, rowIdx, expr.Name) - } - jsonVector := vectorExpr.Expr.(*sqlparser.SQLVal) - golog.Debug("conn", "handleInsert", "jsonVector", c.connectionId, "jsonVector", string(jsonVector.Val)) - err := json.Unmarshal(jsonVector.Val, &values[rowIdx]) - if err != nil { - return nil, errors.Wrap(err, "json.Unmarshal failed") - } - default: - return nil, errors.Errorf("column[%s] row[%d] type[%s] not supported", columnName, rowIdx, reflect.TypeOf(expr).String()) - } - } - dim := len(values[0]) - golog.Debug("conn", "handleInsert", "insert", c.connectionId, "dim", dim) - columns = append(columns, entity.NewColumnFloatVector(columnName, dim, values)) - // TODO: support more types - default: - return nil, errors.Errorf("column[%s] type[%s] not supported", columnName, columnSchema.DataType) + } + } else { + for _, n := range s.Columns { + names = append(names, n.String()) + } + } + seen := map[string]bool{} + for _, n := range names { + f := fields[n] + if f == nil { + return nil, fmt.Errorf("unknown field %s", n) + } + if seen[n] { + return nil, fmt.Errorf("duplicate field %s", n) + } + seen[n] = true + if f.AutoID { + return nil, fmt.Errorf("auto-ID field %s must be omitted", n) + } + } + if s.Action == sqlparser.ReplaceStr { + for _, f := range schema.Fields { + if f.AutoID { + return nil, fmt.Errorf("upsert requires an explicit primary key; auto-ID collections are unsupported") + } + } + } + for _, f := range schema.Fields { + if !f.AutoID && !seen[f.Name] { + return nil, fmt.Errorf("missing field %s", f.Name) + } + } + for i, row := range rows { + if len(row) != len(names) { + return nil, fmt.Errorf("row %d has %d values, expected %d", i+1, len(row), len(names)) + } + } + columns := make([]entity.Column, 0, len(names)) + for j, name := range names { + f := fields[name] + column, err := emptyColumn(f) + if err != nil { + return nil, err + } + for i, row := range rows { + v, err := fieldValue(f, row[j]) + if err != nil { + return nil, fmt.Errorf("field %s row %d: %w", name, i+1, err) + } + if err = column.AppendValue(v); err != nil { + return nil, err } } + columns = append(columns, column) } return columns, nil } +func emptyColumn(f *entity.Field) (entity.Column, error) { + switch f.DataType { + case entity.FieldTypeBool: + return entity.NewColumnBool(f.Name, nil), nil + case entity.FieldTypeInt8: + return entity.NewColumnInt8(f.Name, nil), nil + case entity.FieldTypeInt16: + return entity.NewColumnInt16(f.Name, nil), nil + case entity.FieldTypeInt32: + return entity.NewColumnInt32(f.Name, nil), nil + case entity.FieldTypeInt64: + return entity.NewColumnInt64(f.Name, nil), nil + case entity.FieldTypeFloat: + return entity.NewColumnFloat(f.Name, nil), nil + case entity.FieldTypeDouble: + return entity.NewColumnDouble(f.Name, nil), nil + case entity.FieldTypeVarChar: + return entity.NewColumnVarChar(f.Name, nil), nil + case entity.FieldTypeJSON: + return entity.NewColumnJSONBytes(f.Name, nil), nil + case entity.FieldTypeFloatVector: + dim, e := strconv.Atoi(f.TypeParams["dim"]) + if e != nil || dim <= 0 { + return nil, fmt.Errorf("invalid vector dimension") + } + return entity.NewColumnFloatVector(f.Name, dim, nil), nil + default: + return nil, fmt.Errorf("unsupported field type %s", f.DataType) + } +} +func literalText(e sqlparser.Expr) (string, error) { + switch v := e.(type) { + case *sqlparser.SQLVal: + if v.Type == sqlparser.StrVal || v.Type == sqlparser.IntVal || v.Type == sqlparser.FloatVal { + return string(v.Val), nil + } + case sqlparser.BoolVal: + return strconv.FormatBool(bool(v)), nil + case *sqlparser.UnaryExpr: + if v.Operator == "-" || v.Operator == "+" { + if n, ok := v.Expr.(*sqlparser.SQLVal); ok && (n.Type == sqlparser.IntVal || n.Type == sqlparser.FloatVal) { + return v.Operator + string(n.Val), nil + } + } + } + return "", fmt.Errorf("expected a literal value") +} +func fieldValue(f *entity.Field, e sqlparser.Expr) (interface{}, error) { + if f.DataType == entity.FieldTypeFloatVector { + if fn, ok := e.(*sqlparser.FuncExpr); ok { + if !strings.EqualFold(fn.Name.String(), "json_vector") || len(fn.Exprs) != 1 { + return nil, fmt.Errorf("expected json_vector(string)") + } + a, ok := fn.Exprs[0].(*sqlparser.AliasedExpr) + if !ok { + return nil, fmt.Errorf("invalid vector argument") + } + e = a.Expr + } + } + text, err := literalText(e) + if err != nil { + return nil, err + } + switch f.DataType { + case entity.FieldTypeBool: + if text == "1" { + return true, nil + } + if text == "0" { + return false, nil + } + return strconv.ParseBool(text) + case entity.FieldTypeInt8, entity.FieldTypeInt16, entity.FieldTypeInt32, entity.FieldTypeInt64: + bits := map[entity.FieldType]int{entity.FieldTypeInt8: 8, entity.FieldTypeInt16: 16, entity.FieldTypeInt32: 32, entity.FieldTypeInt64: 64}[f.DataType] + n, e := strconv.ParseInt(text, 10, bits) + if e != nil { + return nil, e + } + switch bits { + case 8: + return int8(n), nil + case 16: + return int16(n), nil + case 32: + return int32(n), nil + } + return n, nil + case entity.FieldTypeFloat, entity.FieldTypeDouble: + bits := 64 + if f.DataType == entity.FieldTypeFloat { + bits = 32 + } + n, e := strconv.ParseFloat(text, bits) + if e != nil { + return nil, e + } + if math.IsNaN(n) || math.IsInf(n, 0) { + return nil, fmt.Errorf("non-finite float") + } + if bits == 32 { + return float32(n), nil + } + return n, nil + case entity.FieldTypeVarChar: + v, ok := e.(*sqlparser.SQLVal) + if !ok || v.Type != sqlparser.StrVal { + return nil, fmt.Errorf("expected string literal") + } + max, err := strconv.Atoi(f.TypeParams["max_length"]) + if err != nil || len(text) > max { + return nil, fmt.Errorf("string exceeds max_length") + } + return text, nil + case entity.FieldTypeJSON: + if !json.Valid([]byte(text)) { + return nil, fmt.Errorf("invalid JSON") + } + return []byte(text), nil + case entity.FieldTypeFloatVector: + var raw []json.RawMessage + if err := json.Unmarshal([]byte(text), &raw); err != nil { + return nil, err + } + for _, x := range raw { + if string(x) == "null" { + return nil, fmt.Errorf("vector elements must be numbers") + } + } + var v []float32 + if err := json.Unmarshal([]byte(text), &v); err != nil { + return nil, err + } + dim, _ := strconv.Atoi(f.TypeParams["dim"]) + if len(v) != dim { + return nil, fmt.Errorf("vector dimension %d, expected %d", len(v), dim) + } + return v, nil + } + return nil, fmt.Errorf("unsupported field type %s", f.DataType) +} +func (c *ClientConn) handleDelete(s *sqlparser.Delete) error { + if len(s.Targets) > 0 || len(s.Partitions) > 0 || s.Limit != nil || len(s.OrderBy) > 0 || s.Where == nil { + return fmt.Errorf("DELETE requires WHERE; multi-table, partition, ORDER BY and LIMIT are unsupported") + } + if len(s.TableExprs) != 1 { + return fmt.Errorf("DELETE requires one collection") + } + a, ok := s.TableExprs[0].(*sqlparser.AliasedTableExpr) + if !ok || !a.As.IsEmpty() || a.Hints != nil { + return fmt.Errorf("unsupported DELETE target") + } + t, ok := a.Expr.(sqlparser.TableName) + if !ok || !t.Qualifier.IsEmpty() { + return fmt.Errorf("unsupported DELETE target") + } + filter, err := scalarFilter(s.Where.Expr) + if err != nil { + return err + } + if err = c.upstream.Delete(c.ctx, t.Name.String(), "", filter); err != nil { + return err + } + // SDK v2 Delete does not expose deleted row count. Do not fabricate one. + return c.writeOK(nil) +} diff --git a/pkg/integration_test.go b/pkg/integration_test.go new file mode 100644 index 0000000..2bd9b43 --- /dev/null +++ b/pkg/integration_test.go @@ -0,0 +1,206 @@ +package pkg + +import ( + "context" + "database/sql" + "fmt" + "os" + "strings" + "testing" + "time" + + _ "github.com/go-sql-driver/mysql" + "github.com/jackc/pgx/v5" +) + +func startTestServer(t *testing.T, cfg *Config) *Server { + t.Helper() + s, err := NewServer(cfg) + if err != nil { + t.Fatal(err) + } + done := make(chan error, 1) + go func() { done <- s.Run() }() + t.Cleanup(func() { + s.Close() + select { + case err := <-done: + if err != nil { + t.Error(err) + } + case <-time.After(5 * time.Second): + t.Error("server shutdown hung") + } + }) + return s +} + +// Enabled explicitly to avoid connecting ordinary unit tests to any Milvus. +func TestMilvusIntegration(t *testing.T) { + addr := os.Getenv("MILVUS_TEST_ADDR") + if addr == "" { + t.Skip("set MILVUS_TEST_ADDR for disposable Milvus integration") + } + s := startTestServer(t, &Config{Mode: "both", Addr: "127.0.0.1:0", PostgresAddr: "127.0.0.1:0", User: "root", Password: "integration", QueryTimeoutSeconds: 90, Milvus: MilvusConfig{Address: addr}}) + ctx, cancel := context.WithTimeout(context.Background(), 3*time.Minute) + defer cancel() + for _, mode := range []string{"mysql", "postgres"} { + t.Run(mode, func(t *testing.T) { + dbname := fmt.Sprintf("sqlproxy_%s_%d", mode, time.Now().UnixNano()) + var exec func(string, ...interface{}) error + var query func(string, ...interface{}) ([][]interface{}, error) + if mode == "mysql" { + db, e := sql.Open("mysql", fmt.Sprintf("root:integration@tcp(%s)/default", s.listeners[0].Addr())) + if e != nil { + t.Fatal(e) + } + defer db.Close() + db.SetMaxOpenConns(1) + exec = func(q string, args ...interface{}) error { _, e := db.ExecContext(ctx, q, args...); return e } + query = func(q string, args ...interface{}) ([][]interface{}, error) { + rs, e := db.QueryContext(ctx, q, args...) + if e != nil { + return nil, e + } + defer rs.Close() + cols, e := rs.Columns() + if e != nil { + return nil, e + } + out := [][]interface{}{} + for rs.Next() { + values := make([]interface{}, len(cols)) + ptrs := make([]interface{}, len(cols)) + for i := range ptrs { + ptrs[i] = &values[i] + } + if e := rs.Scan(ptrs...); e != nil { + return nil, e + } + out = append(out, values) + } + return out, rs.Err() + } + } else { + pg, e := pgx.Connect(ctx, fmt.Sprintf("postgres://root:integration@%s/default?sslmode=disable", s.listeners[1].Addr())) + if e != nil { + t.Fatal(e) + } + defer pg.Close(ctx) + exec = func(q string, args ...interface{}) error { _, e := pg.Exec(ctx, q, args...); return e } + query = func(q string, args ...interface{}) ([][]interface{}, error) { + rs, e := pg.Query(ctx, q, args...) + if e != nil { + return nil, e + } + defer rs.Close() + out := [][]interface{}{} + for rs.Next() { + v, e := rs.Values() + if e != nil { + return nil, e + } + out = append(out, v) + } + return out, rs.Err() + } + } + mustExec := func(q string, args ...interface{}) { + t.Helper() + if e := exec(q, args...); e != nil { + t.Fatalf("%s: %v", q, e) + } + } + mustQuery := func(q string, args ...interface{}) [][]interface{} { + t.Helper() + v, e := query(q, args...) + if e != nil { + t.Fatalf("%s: %v", q, e) + } + return v + } + mustExec("CREATE DATABASE " + dbname) + defer func() { + exec("USE default") + if e := exec("DROP DATABASE " + dbname); e != nil { + t.Error(e) + } + }() + mustExec("USE " + dbname) + mustExec("CREATE TABLE items (id bigint PRIMARY KEY, name varchar(100), enabled bool, score double, meta json, embedding vector(3))") + mustExec("CREATE INDEX embedding_idx ON items (embedding) USING FLAT WITH (metric_type='L2')") + mustExec("INSERT INTO items VALUES (1,'one',true,1.5,'{\"tag\":1}',json_vector('[1,0,0]')), (2,'two',false,2.5,'{\"tag\":2}',json_vector('[0,1,0]'))") + mustExec("LOAD TABLE items") + if v := mustQuery("SHOW INDEXES FROM items"); len(v) != 1 { + t.Fatalf("indexes: %v", v) + } + if v := mustQuery("DESCRIBE items"); len(v) != 6 { + t.Fatalf("describe: %v", v) + } + if v := mustQuery("SELECT id,name,enabled,score,meta,embedding FROM items WHERE id=1"); len(v) != 1 { + t.Fatalf("rows: %v", v) + } + if v := mustQuery("SELECT id FROM items WHERE id=999"); len(v) != 0 { + t.Fatal(v) + } + if v := mustQuery("SELECT count(*) FROM items"); len(v) != 1 { + t.Fatal(v) + } + if v := mustQuery("SELECT id,_distance FROM items WHERE embedding LIKE json_vector('[1,0,0]') LIMIT 1"); len(v) != 1 || fmt.Sprint(v[0][0]) != "1" { + t.Fatalf("search: %v", v) + } + placeholder := "?" + if mode == "postgres" { + placeholder = "$1" + } + if v := mustQuery("SELECT id,name FROM items WHERE id="+placeholder, int64(2)); len(v) != 1 { + t.Fatalf("parameter query: %v", v) + } + mustExec("UPSERT INTO items VALUES (2,'updated',true,3.5,'{}',json_vector('[0,0,1]'))") + if v := mustQuery("SELECT name FROM items WHERE id=2"); len(v) != 1 { + t.Fatal(v) + } + mustExec("DELETE FROM items WHERE id=2") + if v := mustQuery("SELECT id FROM items WHERE id=2"); len(v) != 0 { + t.Fatalf("delete: %v", v) + } + for _, bad := range []string{"DROP VIEW items", "SELECT 1 HAVING 1=0", "DELETE FROM items", "INSERT INTO items VALUES (3,'bad',true,1,'{}',json_vector('[null,0,1]'))"} { + if e := exec(bad); e == nil { + t.Fatalf("accepted unsupported SQL %s", bad) + } + } + mustQuery("SELECT id FROM items WHERE id=1") + mustExec("FLUSH TABLE items") + mustExec("RELEASE TABLE items") + mustExec("DROP INDEX embedding_idx ON items") + mustExec("CREATE PARTITION extra ON items") + if v := mustQuery("SHOW PARTITIONS FROM items"); len(v) != 2 { + t.Fatal(v) + } + mustExec("DROP PARTITION extra ON items") + mustExec("DROP TABLE items") + }) + } + // Incorrect credentials must fail both protocol handshakes. + for i, mode := range []string{"mysql", "postgres"} { + if mode == "mysql" { + db, e := sql.Open("mysql", fmt.Sprintf("root:wrong@tcp(%s)/default", s.listeners[i].Addr())) + if e != nil { + t.Fatal(e) + } + if e = db.PingContext(ctx); e == nil { + t.Fatal("MySQL accepted wrong password") + } + db.Close() + } else { + pg, e := pgx.Connect(ctx, fmt.Sprintf("postgres://root:wrong@%s/default?sslmode=disable", s.listeners[i].Addr())) + if e == nil { + pg.Close(ctx) + t.Fatal("PG accepted wrong password") + } + if !strings.Contains(e.Error(), "authentication") { + t.Fatal(e) + } + } + } +} diff --git a/pkg/mysql.go b/pkg/mysql.go new file mode 100644 index 0000000..6e1b9b9 --- /dev/null +++ b/pkg/mysql.go @@ -0,0 +1,130 @@ +package pkg + +import ( + "context" + "encoding/json" + "fmt" + "net" + "time" + + legacy "github.com/flike/kingshard/mysql" + "github.com/go-mysql-org/go-mysql/mysql" + "github.com/go-mysql-org/go-mysql/server" +) + +type mysqlHandler struct { + ctx context.Context + session *ClientConn + timeout time.Duration +} + +func (s *Server) serveMySQL(ctx context.Context, co net.Conn) { + startup, cancel := context.WithTimeout(ctx, 15*time.Second) + session, err := s.newSession(startup) + cancel() + if err != nil { + return + } + defer session.Close() + h := &mysqlHandler{session: session, ctx: ctx, timeout: time.Duration(s.cfg.QueryTimeoutSeconds) * time.Second} + conf := server.NewServer("8.0.11-milvus-sql-proxy", mysql.DEFAULT_COLLATION_ID, mysql.AUTH_NATIVE_PASSWORD, nil, s.tlsConfig) + conn, err := conf.NewConn(co, s.cfg.User, s.cfg.Password, h) + if err != nil { + return + } + defer conn.Close() + co.SetDeadline(time.Time{}) + for { + if err := conn.HandleCommand(); err != nil { + return + } + } +} +func (h *mysqlHandler) UseDB(db string) error { + ctx, cancel := context.WithTimeout(h.ctx, h.timeout) + defer cancel() + h.session.ctx = ctx + return h.session.handleUseDB(db, nil) +} +func (h *mysqlHandler) HandleQuery(q string) (*mysql.Result, error) { return h.run(q, false) } +func (h *mysqlHandler) run(q string, binary bool) (*mysql.Result, error) { + ctx, cancel := context.WithTimeout(h.ctx, h.timeout) + defer cancel() + r, err := h.session.Execute(ctx, q) + if err != nil { + return nil, err + } + return mysqlResult(r, binary) +} +func mysqlResult(r *legacy.Result, binary bool) (*mysql.Result, error) { + out := &mysql.Result{AffectedRows: r.AffectedRows, InsertId: r.InsertId, Status: mysql.SERVER_STATUS_AUTOCOMMIT} + if r.Resultset == nil { + return out, nil + } + names := make([]string, len(r.Fields)) + for i, f := range r.Fields { + names[i] = string(f.Name) + } + values := make([][]interface{}, len(r.Values)) + for i, row := range r.Values { + values[i] = make([]interface{}, len(row)) + for j, v := range row { + switch x := v.(type) { + case bool: + if x { + v = int64(1) + } else { + v = int64(0) + } + case []float32: + b, _ := json.Marshal(x) + v = string(b) + } + values[i][j] = v + } + } + rs, err := mysql.BuildSimpleResultset(names, values, binary) + if err != nil { + return nil, err + } + // Preserve field metadata for empty result sets as well as nonempty rows. + for i, f := range r.Fields { + if len(values) == 0 { + rs.Fields[i].Type = f.Type + rs.Fields[i].Charset = f.Charset + } + } + out.Resultset = rs + return out, nil +} +func (h *mysqlHandler) HandleFieldList(string, string) ([]*mysql.Field, error) { + return nil, fmt.Errorf("COM_FIELD_LIST is not supported; use DESCRIBE table") +} +func (h *mysqlHandler) HandleStmtPrepare(q string) (int, int, interface{}, error) { + normalized, n, err := bindSQL(q, "mysql", nil, true) + if err != nil { + return 0, 0, nil, err + } + ctx, cancel := context.WithTimeout(h.ctx, h.timeout) + defer cancel() + r, err := h.session.Describe(ctx, normalized) + if err != nil { + return 0, 0, nil, err + } + cols := 0 + if r.Resultset != nil { + cols = len(r.Fields) + } + return n, cols, nil, nil +} +func (h *mysqlHandler) HandleStmtExecute(_ interface{}, q string, args []interface{}) (*mysql.Result, error) { + q, _, err := bindSQL(q, "mysql", args, false) + if err != nil { + return nil, err + } + return h.run(q, true) +} +func (h *mysqlHandler) HandleStmtClose(interface{}) error { return nil } +func (h *mysqlHandler) HandleOtherCommand(byte, []byte) error { + return fmt.Errorf("unsupported MySQL command") +} diff --git a/pkg/parameters.go b/pkg/parameters.go new file mode 100644 index 0000000..204722b --- /dev/null +++ b/pkg/parameters.go @@ -0,0 +1,150 @@ +package pkg + +import ( + "context" + "fmt" + "strconv" + "strings" + + "github.com/milvus-io/milvus-sdk-go/v2/entity" + "github.com/xwb1989/sqlparser" +) + +func (c *ClientConn) parameterOIDs(ctx context.Context, q string, oids []uint32) ([]uint32, error) { + if len(oids) == 0 { + return oids, nil + } + args := make([]interface{}, len(oids)) + for i := range args { + args[i] = fmt.Sprintf("__milvus_param_%d__", i+1) + } + normalized, _, err := bindSQL(q, "postgres", args, false) + if err != nil { + return nil, err + } + normalized = upsertRE.ReplaceAllString(normalized, "REPLACE INTO") + stmt, err := sqlparser.Parse(normalized) + if err != nil { + for i := range oids { + if oids[i] == 0 { + oids[i] = 25 + } + } + return oids, nil + } + var table string + switch s := stmt.(type) { + case *sqlparser.Select: + if len(s.From) == 1 { + if a, ok := s.From[0].(*sqlparser.AliasedTableExpr); ok { + if t, ok := a.Expr.(sqlparser.TableName); ok { + table = t.Name.String() + } + } + } + case *sqlparser.Insert: + table = s.Table.Name.String() + case *sqlparser.Delete: + if len(s.TableExprs) == 1 { + if a, ok := s.TableExprs[0].(*sqlparser.AliasedTableExpr); ok { + if t, ok := a.Expr.(sqlparser.TableName); ok { + table = t.Name.String() + } + } + } + } + fields := map[string]*entity.Field{} + var schema *entity.Schema + if table != "" && table != "dual" { + old := c.ctx + c.ctx = ctx + schema, err = c.GetCollectinSchema(table) + c.ctx = old + if err != nil { + return nil, err + } + for _, f := range schema.Fields { + fields[f.Name] = f + } + } + set := func(e sqlparser.Expr, oid uint32) { + v, ok := e.(*sqlparser.SQLVal) + if !ok { + return + } + text := string(v.Val) + if !strings.HasPrefix(text, "__milvus_param_") || !strings.HasSuffix(text, "__") { + return + } + n, err := strconv.Atoi(strings.TrimSuffix(strings.TrimPrefix(text, "__milvus_param_"), "__")) + if err == nil && n > 0 && n <= len(oids) && oids[n-1] == 0 { + oids[n-1] = oid + } + } + fieldOID := func(f *entity.Field) uint32 { + if f == nil { + return 25 + } + return pgOID(resultField(f.Name, f.DataType)) + } + err = sqlparser.Walk(func(node sqlparser.SQLNode) (bool, error) { + switch x := node.(type) { + case *sqlparser.ComparisonExpr: + if x == nil { + return true, nil + } + if col, ok := x.Left.(*sqlparser.ColName); ok { + oid := fieldOID(fields[col.Name.String()]) + set(x.Right, oid) + if tuple, ok := x.Right.(sqlparser.ValTuple); ok { + for _, e := range tuple { + set(e, oid) + } + } + } + if col, ok := x.Right.(*sqlparser.ColName); ok { + set(x.Left, fieldOID(fields[col.Name.String()])) + } + case *sqlparser.Limit: + if x == nil { + return true, nil + } + set(x.Rowcount, 20) + if x.Offset != nil { + set(x.Offset, 20) + } + } + return true, nil + }, stmt) + if err != nil { + return nil, err + } + if ins, ok := stmt.(*sqlparser.Insert); ok && schema != nil { + names := []string{} + for _, n := range ins.Columns { + names = append(names, n.String()) + } + if len(names) == 0 { + for _, f := range schema.Fields { + if !f.AutoID { + names = append(names, f.Name) + } + } + } + if rows, ok := ins.Rows.(sqlparser.Values); ok { + for _, row := range rows { + for j, e := range row { + if j < len(names) { + set(e, fieldOID(fields[names[j]])) + } + } + } + } + } + for i := range oids { + if oids[i] == 0 { + oids[i] = 25 + } + } + return oids, nil +} diff --git a/pkg/postgres.go b/pkg/postgres.go new file mode 100644 index 0000000..b90248a --- /dev/null +++ b/pkg/postgres.go @@ -0,0 +1,501 @@ +package pkg + +import ( + "context" + "crypto/subtle" + "crypto/tls" + "encoding/binary" + "fmt" + "math" + "net" + "strconv" + "strings" + "time" + + legacy "github.com/flike/kingshard/mysql" + "github.com/jackc/pgx/v5/pgproto3" +) + +type pgStatement struct { + query string + oids []uint32 +} +type pgPortal struct { + query string + formats []int16 + result *legacy.Result + pos int + completed bool +} + +func (s *Server) servePostgres(ctx context.Context, co net.Conn) { + b := pgproto3.NewBackend(co, co) + b.SetMaxBodyLen(16 << 20) + startup, err := b.ReceiveStartupMessage() + if err != nil { + return + } + if _, ok := startup.(*pgproto3.SSLRequest); ok { + if s.tlsConfig == nil { + if _, err = co.Write([]byte("N")); err != nil { + return + } + } else { + if _, err = co.Write([]byte("S")); err != nil { + return + } + t := tls.Server(co, s.tlsConfig) + if err = t.HandshakeContext(ctx); err != nil { + return + } + co = t + b = pgproto3.NewBackend(co, co) + b.SetMaxBodyLen(16 << 20) + } + startup, err = b.ReceiveStartupMessage() + if err != nil { + return + } + } + start, ok := startup.(*pgproto3.StartupMessage) + if !ok { + return + } + fail := func(e error) { b.Send(&pgproto3.ErrorResponse{Severity: "ERROR", Code: "0A000", Message: e.Error()}) } + b.Send(&pgproto3.AuthenticationCleartextPassword{}) + if b.Flush() != nil { + return + } + auth, e := b.Receive() + if e != nil { + return + } + pw, ok := auth.(*pgproto3.PasswordMessage) + if !ok || subtle.ConstantTimeCompare([]byte(start.Parameters["user"]), []byte(s.cfg.User)) != 1 || subtle.ConstantTimeCompare([]byte(pw.Password), []byte(s.cfg.Password)) != 1 { + b.Send(&pgproto3.ErrorResponse{Severity: "FATAL", Code: "28P01", Message: "authentication failed"}) + b.Flush() + return + } + startupCtx, startupCancel := context.WithTimeout(ctx, 15*time.Second) + defer startupCancel() + session, err := s.newSession(startupCtx) + if err != nil { + fail(err) + b.Flush() + return + } + defer session.Close() + db := start.Parameters["database"] + if db == "" { + db = "default" + } + session.ctx = startupCtx + if err = session.handleUseDB(db, nil); err != nil { + fail(err) + b.Flush() + return + } + b.Send(&pgproto3.AuthenticationOk{}) + for _, p := range [][2]string{{"server_version", "14.0"}, {"client_encoding", "UTF8"}, {"standard_conforming_strings", "on"}, {"DateStyle", "ISO, MDY"}, {"TimeZone", "UTC"}} { + b.Send(&pgproto3.ParameterStatus{Name: p[0], Value: p[1]}) + } + b.Send(&pgproto3.ReadyForQuery{TxStatus: 'I'}) + if b.Flush() != nil { + return + } + co.SetDeadline(time.Time{}) + statements := map[string]pgStatement{} + portals := map[string]*pgPortal{} + failed := false + timeout := time.Duration(s.cfg.QueryTimeoutSeconds) * time.Second + execute := func(q string, describe bool) (*legacy.Result, error) { + qctx, cancel := context.WithTimeout(ctx, timeout) + defer cancel() + if describe { + return session.Describe(qctx, q) + } + return session.Execute(qctx, q) + } + for { + msg, err := b.Receive() + if err != nil { + return + } + if _, ok := msg.(*pgproto3.Terminate); ok { + return + } + if failed { + if _, ok := msg.(*pgproto3.Sync); !ok { + continue + } + } + var opErr error + switch m := msg.(type) { + case *pgproto3.Query: + if strings.TrimSpace(m.String) == "" { + b.Send(&pgproto3.EmptyQueryResponse{}) + } else { + q, _, e := bindSQL(m.String, "postgres", nil, false) + if e != nil { + opErr = e + } else { + r, e := execute(q, false) + if e != nil { + opErr = e + } else { + if r.Resultset != nil { + b.Send(pgDescription(r, nil)) + opErr = pgRows(b, r, 0, len(r.Values), nil) + } + if opErr == nil { + b.Send(&pgproto3.CommandComplete{CommandTag: []byte(pgTag(q, r))}) + } + } + } + } + if opErr != nil { + fail(opErr) + opErr = nil + } + b.Send(&pgproto3.ReadyForQuery{TxStatus: 'I'}) + case *pgproto3.Parse: + if len(statements) >= 1024 && m.Name != "" { + opErr = fmt.Errorf("too many prepared statements") + break + } + _, n, e := bindSQL(m.Query, "postgres", nil, true) + if e != nil { + opErr = e + break + } + if len(m.ParameterOIDs) > n { + opErr = fmt.Errorf("too many parameter types") + break + } + oids := make([]uint32, n) + copy(oids, m.ParameterOIDs) + parseCtx, parseCancel := context.WithTimeout(ctx, timeout) + oids, e = session.parameterOIDs(parseCtx, m.Query, oids) + parseCancel() + if e != nil { + opErr = e + break + } + if _, exists := statements[m.Name]; exists && m.Name != "" { + opErr = fmt.Errorf("prepared statement already exists") + break + } + statements[m.Name] = pgStatement{m.Query, oids} + b.Send(&pgproto3.ParseComplete{}) + case *pgproto3.Describe: + var q string + var formats []int16 + if m.ObjectType == 'S' { + st, exists := statements[m.Name] + if !exists { + opErr = fmt.Errorf("unknown prepared statement") + break + } + b.Send(&pgproto3.ParameterDescription{ParameterOIDs: st.oids}) + q, _, opErr = bindSQL(st.query, "postgres", nil, true) + } else if m.ObjectType == 'P' { + p, exists := portals[m.Name] + if !exists { + opErr = fmt.Errorf("unknown portal") + break + } + q = p.query + formats = p.formats + } else { + opErr = fmt.Errorf("invalid describe target") + break + } + if opErr == nil { + r, e := execute(q, true) + if e != nil { + opErr = e + } else if r.Resultset == nil { + b.Send(&pgproto3.NoData{}) + } else { + b.Send(pgDescription(r, formats)) + } + } + case *pgproto3.Bind: + st, exists := statements[m.PreparedStatement] + if !exists { + opErr = fmt.Errorf("unknown prepared statement") + break + } + if len(m.Parameters) != len(st.oids) { + opErr = fmt.Errorf("parameter count mismatch") + break + } + if !validFormats(m.ParameterFormatCodes, len(m.Parameters)) { + opErr = fmt.Errorf("invalid parameter formats") + break + } + args := make([]interface{}, len(m.Parameters)) + for i, data := range m.Parameters { + args[i], opErr = pgParameter(data, st.oids[i], formatAt(m.ParameterFormatCodes, i)) + if opErr != nil { + break + } + } + if opErr != nil { + break + } + q, _, e := bindSQL(st.query, "postgres", args, false) + if e != nil { + opErr = e + break + } + r, e := execute(q, true) + if e != nil { + opErr = e + break + } + cols := 0 + if r.Resultset != nil { + cols = len(r.Fields) + } + if !validFormats(m.ResultFormatCodes, cols) { + opErr = fmt.Errorf("invalid result formats") + break + } + if len(portals) >= 1024 && m.DestinationPortal != "" { + opErr = fmt.Errorf("too many portals") + break + } + if _, exists := portals[m.DestinationPortal]; exists && m.DestinationPortal != "" { + opErr = fmt.Errorf("portal already exists") + break + } + portals[m.DestinationPortal] = &pgPortal{query: q, formats: append([]int16(nil), m.ResultFormatCodes...)} + b.Send(&pgproto3.BindComplete{}) + case *pgproto3.Execute: + p, exists := portals[m.Portal] + if !exists { + opErr = fmt.Errorf("unknown portal") + break + } + if p.result == nil { + p.result, opErr = execute(p.query, false) + if opErr != nil { + break + } + } + r := p.result + if r.Resultset != nil && !p.completed { + end := len(r.Values) + if m.MaxRows > 0 && int(m.MaxRows) < end-p.pos { + end = p.pos + int(m.MaxRows) + } + opErr = pgRows(b, r, p.pos, end, p.formats) + p.pos = end + if opErr != nil { + break + } + if end < len(r.Values) { + b.Send(&pgproto3.PortalSuspended{}) + break + } + } + p.completed = true + b.Send(&pgproto3.CommandComplete{CommandTag: []byte(pgTag(p.query, r))}) + case *pgproto3.Close: + switch m.ObjectType { + case 'S': + delete(statements, m.Name) + case 'P': + delete(portals, m.Name) + default: + opErr = fmt.Errorf("invalid close target") + } + if opErr == nil { + b.Send(&pgproto3.CloseComplete{}) + } + case *pgproto3.Sync: + failed = false + portals = map[string]*pgPortal{} + b.Send(&pgproto3.ReadyForQuery{TxStatus: 'I'}) + case *pgproto3.Flush: + default: + opErr = fmt.Errorf("unsupported PostgreSQL message %T", msg) + } + if opErr != nil { + fail(opErr) + failed = true + } + if b.Flush() != nil { + return + } + } +} +func formatAt(f []int16, i int) int16 { + if len(f) == 0 { + return 0 + } + if len(f) == 1 { + return f[0] + } + if i >= len(f) { + return 0 + } + return f[i] +} +func validFormats(f []int16, n int) bool { + if len(f) != 0 && len(f) != 1 && len(f) != n { + return false + } + for _, v := range f { + if v != 0 && v != 1 { + return false + } + } + return true +} +func pgOID(f *legacy.Field) uint32 { + switch f.Type { + case legacy.MYSQL_TYPE_TINY: + return 16 + case legacy.MYSQL_TYPE_LONGLONG: + return 20 + case legacy.MYSQL_TYPE_DOUBLE: + return 701 + default: + return 25 + } +} +func pgDescription(r *legacy.Result, formats []int16) *pgproto3.RowDescription { + d := &pgproto3.RowDescription{} + for i, f := range r.Fields { + oid := pgOID(f) + size := int16(-1) + if oid == 20 || oid == 701 { + size = 8 + } + if oid == 16 { + size = 1 + } + d.Fields = append(d.Fields, pgproto3.FieldDescription{Name: f.Name, DataTypeOID: oid, DataTypeSize: size, TypeModifier: -1, Format: formatAt(formats, i)}) + } + return d +} +func pgRows(b *pgproto3.Backend, r *legacy.Result, start, end int, formats []int16) error { + for _, row := range r.Values[start:end] { + values := make([][]byte, len(row)) + for i, v := range row { + var err error + values[i], err = pgValue(v, pgOID(r.Fields[i]), formatAt(formats, i)) + if err != nil { + return err + } + } + b.Send(&pgproto3.DataRow{Values: values}) + } + return nil +} +func pgValue(v interface{}, oid uint32, format int16) ([]byte, error) { + if v == nil { + return nil, nil + } + if format == 0 { + return formatValue(v) + } + switch oid { + case 16: + x, ok := v.(bool) + if !ok { + return nil, fmt.Errorf("expected bool") + } + if x { + return []byte{1}, nil + } + return []byte{0}, nil + case 20: + n, e := strconv.ParseInt(fmt.Sprint(v), 10, 64) + if e != nil { + return nil, e + } + b := make([]byte, 8) + binary.BigEndian.PutUint64(b, uint64(n)) + return b, nil + case 701: + n, e := strconv.ParseFloat(fmt.Sprint(v), 64) + if e != nil { + return nil, e + } + b := make([]byte, 8) + binary.BigEndian.PutUint64(b, math.Float64bits(n)) + return b, nil + default: + return formatValue(v) + } +} +func pgParameter(data []byte, oid uint32, format int16) (interface{}, error) { + if data == nil { + return nil, nil + } + if format == 0 { + switch oid { + case 16: + return strconv.ParseBool(string(data)) + case 20, 21, 23: + return strconv.ParseInt(string(data), 10, 64) + case 700, 701: + return strconv.ParseFloat(string(data), 64) + default: + return string(data), nil + } + } + switch oid { + case 16: + if len(data) == 1 && data[0] <= 1 { + return data[0] == 1, nil + } + case 20: + if len(data) == 8 { + return int64(binary.BigEndian.Uint64(data)), nil + } + case 23: + if len(data) == 4 { + return int32(binary.BigEndian.Uint32(data)), nil + } + case 21: + if len(data) == 2 { + return int16(binary.BigEndian.Uint16(data)), nil + } + case 701: + if len(data) == 8 { + return math.Float64frombits(binary.BigEndian.Uint64(data)), nil + } + case 700: + if len(data) == 4 { + return math.Float32frombits(binary.BigEndian.Uint32(data)), nil + } + case 25, 1043, 114: + return string(data), nil + } + return nil, fmt.Errorf("unsupported binary parameter OID %d or invalid length", oid) +} +func pgTag(q string, r *legacy.Result) string { + words := strings.Fields(q) + if len(words) == 0 { + return "" + } + verb := strings.ToUpper(words[0]) + if r.Resultset != nil { + return fmt.Sprintf("SELECT %d", len(r.Values)) + } + switch verb { + case "INSERT", "UPSERT", "REPLACE": + return fmt.Sprintf("INSERT 0 %d", r.AffectedRows) + case "DELETE": + return fmt.Sprintf("DELETE %d", r.AffectedRows) + case "CREATE", "DROP": + if len(words) > 1 { + return verb + " " + strings.ToUpper(words[1]) + } + } + return verb +} diff --git a/pkg/protocol_test.go b/pkg/protocol_test.go new file mode 100644 index 0000000..6f94b49 --- /dev/null +++ b/pkg/protocol_test.go @@ -0,0 +1,285 @@ +package pkg + +import ( + "context" + "database/sql" + "encoding/binary" + "fmt" + "net" + "reflect" + "strings" + "testing" + "time" + + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgproto3" +) + +func mockServer(t *testing.T) *Server { + t.Helper() + s, e := NewServer(&Config{Mode: "both", Addr: "127.0.0.1:0", PostgresAddr: "127.0.0.1:0", User: "root", Password: "secret"}) + if e != nil { + t.Fatal(e) + } + s.newSession = func(ctx context.Context) (*ClientConn, error) { + return NewSession(ctx, &operationMock{schema: testSchema(), db: "default"}), nil + } + done := make(chan error, 1) + go func() { done <- s.Run() }() + t.Cleanup(func() { + s.Close() + select { + case e := <-done: + if e != nil { + t.Error(e) + } + case <-time.After(5 * time.Second): + t.Error("shutdown hung") + } + }) + return s +} +func TestMySQLDriver(t *testing.T) { + s := mockServer(t) + db, e := sql.Open("mysql", fmt.Sprintf("root:secret@tcp(%s)/otherdb", s.listeners[0].Addr())) + if e != nil { + t.Fatal(e) + } + defer db.Close() + db.SetMaxOpenConns(1) + var name string + var id int64 + var flag bool + var score float64 + if e = db.QueryRow("SELECT name,id,enabled,score FROM items WHERE id=?", int64(7)).Scan(&name, &id, &flag, &score); e != nil { + t.Fatal(e) + } + if name != "one" || id != 7 || !flag || score != 1.5 { + t.Fatal(name, id, flag, score) + } + if e = db.QueryRow("SELECT database()").Scan(&name); e != nil || name != "otherdb" { + t.Fatal(name, e) + } + if _, e = db.Exec("BEGIN"); e == nil { + t.Fatal("accepted transaction") + } + if e = db.Ping(); e != nil { + t.Fatal(e) + } + bad, _ := sql.Open("mysql", fmt.Sprintf("root:bad@tcp(%s)/default", s.listeners[0].Addr())) + defer bad.Close() + if e = bad.Ping(); e == nil { + t.Fatal("accepted bad password") + } +} +func TestPostgresDriver(t *testing.T) { + s := mockServer(t) + ctx := context.Background() + pg, e := pgx.Connect(ctx, fmt.Sprintf("postgres://root:secret@%s/otherdb?sslmode=disable", s.listeners[1].Addr())) + if e != nil { + t.Fatal(e) + } + defer pg.Close(ctx) + var name string + var id int64 + var flag bool + var score float64 + if e = pg.QueryRow(ctx, "SELECT name,id,enabled,score FROM items WHERE id=$1", int64(7)).Scan(&name, &id, &flag, &score); e != nil { + t.Fatal(e) + } + if name != "one" || id != 7 || !flag || score != 1.5 { + t.Fatal(name, id, flag, score) + } + if e = pg.QueryRow(ctx, "SELECT current_database()").Scan(&name); e != nil || name != "otherdb" { + t.Fatal(name, e) + } + rows, e := pg.Query(ctx, "DESCRIBE items") + if e != nil { + t.Fatal(e) + } + count := 0 + for rows.Next() { + var field, kind, params string + var primary, auto bool + if e := rows.Scan(&field, &kind, &primary, &auto, ¶ms); e != nil { + t.Fatal(e) + } + count++ + } + if e := rows.Err(); e != nil { + t.Fatal(e) + } + rows.Close() + if count != 6 { + t.Fatal(count) + } + if _, e = pg.Exec(ctx, "BEGIN"); e == nil { + t.Fatal("accepted transaction") + } + if e = pg.QueryRow(ctx, "SELECT 1").Scan(&id); e != nil || id != 1 { + t.Fatal(id, e) + } + // Simple protocol uses PostgreSQL identifier/literal quoting too. + if e = pg.QueryRow(ctx, `SELECT "id" FROM items WHERE name='it''s \\ literal'`, pgx.QueryExecModeSimpleProtocol).Scan(&id); e != nil { + t.Fatal(e) + } + bad, e := pgx.Connect(ctx, fmt.Sprintf("postgres://root:bad@%s/default?sslmode=disable", s.listeners[1].Addr())) + if e == nil { + bad.Close(ctx) + t.Fatal("accepted bad password") + } +} +func TestBindSQL(t *testing.T) { + for _, tc := range []struct { + q, d string + args []interface{} + want string + n int + }{ + {"select ?,'?',`?`,\"?\" -- ?\n from x where id=?", "mysql", []interface{}{int64(1), "a' OR 1=1 --"}, "select 1,'?',`?`,\"?\" -- ?\n from x where id='a\\' OR 1=1 --'", 2}, + {`select "id" from t where id=$1 or id=$1 /* $2 */`, "postgres", []interface{}{int64(7)}, "select `id` from t where id=7 or id=7 /* $2 */", 1}, + {`select 'it''s' from t where x=$1`, "postgres", []interface{}{`a\b`}, `select 'it\'s' from t where x='a\\b'`, 1}, + } { + got, n, e := bindSQL(tc.q, tc.d, tc.args, false) + if e != nil || got != tc.want || n != tc.n { + t.Fatalf("%q => %q %d %v", tc.q, got, n, e) + } + } + for _, q := range []string{"select $0", "select $99999999999999999", "select $$str$$", "select 'oops", "select /*oops"} { + if _, _, e := bindSQL(q, "postgres", nil, false); e == nil { + t.Fatal(q) + } + } + if _, _, e := bindSQL("select ?", "mysql", nil, false); e == nil { + t.Fatal("missing arg accepted") + } + if _, _, e := bindSQL("select 1", "mysql", []interface{}{1}, false); e == nil { + t.Fatal("extra arg accepted") + } +} +func TestPGEncodings(t *testing.T) { + for _, tc := range []struct { + oid uint32 + v interface{} + }{{16, true}, {16, false}, {20, int64(-9)}, {701, float64(1.25)}, {25, "str"}} { + b, e := pgValue(tc.v, tc.oid, 1) + if e != nil { + t.Fatal(e) + } + v, e := pgParameter(b, tc.oid, 1) + if e != nil || !reflect.DeepEqual(v, tc.v) { + t.Fatal(v, e) + } + } + b := make([]byte, 4) + binary.BigEndian.PutUint32(b, 42) + if v, e := pgParameter(b, 23, 1); e != nil || v != int32(42) { + t.Fatal(v, e) + } + if _, e := pgParameter([]byte{1}, 20, 1); e == nil { + t.Fatal("invalid int accepted") + } + if v, e := pgParameter(nil, 25, 0); e != nil || v != nil { + t.Fatal(v, e) + } + if validFormats([]int16{2}, 1) || validFormats([]int16{0, 1}, 3) { + t.Fatal("bad formats accepted") + } +} +func TestShutdownCancelsDial(t *testing.T) { + s, e := NewServer(&Config{Addr: "127.0.0.1:0"}) + if e != nil { + t.Fatal(e) + } + entered := make(chan struct{}) + s.newSession = func(ctx context.Context) (*ClientConn, error) { close(entered); <-ctx.Done(); return nil, ctx.Err() } + done := make(chan error, 1) + go func() { done <- s.Run() }() + co, e := net.Dial("tcp", s.listeners[0].Addr().String()) + if e != nil { + t.Fatal(e) + } + defer co.Close() + <-entered + s.Close() + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("shutdown did not cancel upstream dial") + } +} +func TestConfig(t *testing.T) { + for _, data := range []string{"mode: unsupported", "queryTimeoutSeconds: -1", "tlsCert: file"} { + if _, e := ParseConfigData([]byte(data)); e == nil { + t.Fatal(data) + } + } + cfg, e := ParseConfigData([]byte("mode: both")) + if e != nil || cfg.User != "root" || !strings.HasPrefix(cfg.Addr, "127.0.0.1") { + t.Fatal(cfg, e) + } +} + +// Low-level extended protocol regression: errors recover at Sync and a failed +// mutation is never executed during Parse/Describe/Bind. +func TestPostgresExtendedRecovery(t *testing.T) { + s := mockServer(t) + co, e := net.Dial("tcp", s.listeners[1].Addr().String()) + if e != nil { + t.Fatal(e) + } + defer co.Close() + co.SetDeadline(time.Now().Add(5 * time.Second)) + f := pgproto3.NewFrontend(co, co) + f.Send(&pgproto3.StartupMessage{ProtocolVersion: 196608, Parameters: map[string]string{"user": "root", "database": "default"}}) + f.Flush() + if _, e = f.Receive(); e != nil { + t.Fatal(e) + } + f.Send(&pgproto3.PasswordMessage{Password: "secret"}) + f.Flush() + for { + m, e := f.Receive() + if e != nil { + t.Fatal(e) + } + if _, ok := m.(*pgproto3.ReadyForQuery); ok { + break + } + } + f.Send(&pgproto3.Parse{Name: "bad", Query: "SELECT id FROM items WHERE id=$1"}) + f.Send(&pgproto3.Bind{PreparedStatement: "bad", Parameters: [][]byte{[]byte("invalid int")}}) + f.Send(&pgproto3.Execute{}) + f.Send(&pgproto3.Sync{}) + f.Flush() + sawError := false + for { + m, e := f.Receive() + if e != nil { + t.Fatal(e) + } + if _, ok := m.(*pgproto3.ErrorResponse); ok { + sawError = true + } + if _, ok := m.(*pgproto3.ReadyForQuery); ok { + break + } + } + if !sawError { + t.Fatal("expected bind error") + } + f.Send(&pgproto3.Query{String: "SELECT 1"}) + f.Flush() + for { + m, e := f.Receive() + if e != nil { + t.Fatal(e) + } + if _, ok := m.(*pgproto3.ErrorResponse); ok { + t.Fatal(m) + } + if _, ok := m.(*pgproto3.ReadyForQuery); ok { + break + } + } +} diff --git a/pkg/select.go b/pkg/select.go index bb11d10..da2a226 100644 --- a/pkg/select.go +++ b/pkg/select.go @@ -1,78 +1,220 @@ package pkg import ( - "github.com/flike/kingshard/core/golog" + "fmt" + "strings" + "github.com/flike/kingshard/mysql" "github.com/milvus-io/milvus-sdk-go/v2/client" - "github.com/pkg/errors" + "github.com/milvus-io/milvus-sdk-go/v2/entity" "github.com/xwb1989/sqlparser" ) -func (c *ClientConn) handleSelect(stmt *sqlparser.Select, args []interface{}) error { - golog.Debug("conn", "handleSelect", "select", c.connectionId, "stmt", stmt, "len(froms)", len(stmt.From)) - froms := stmt.From - if len(froms) > 1 { - // TODO: support join? - err := errors.Errorf("select from more than one table not supported") - return c.writeError(err) +func (c *ClientConn) handleSelect(stmt *sqlparser.Select, _ []interface{}) error { + if len(stmt.From) == 1 && sqlparser.String(stmt.From[0]) == "dual" && stmt.Where == nil && stmt.Limit == nil && stmt.Distinct == "" && len(stmt.GroupBy) == 0 && stmt.Having == nil && len(stmt.OrderBy) == 0 && stmt.Lock == "" && len(stmt.SelectExprs) == 1 { + expr := strings.ToLower(sqlparser.String(stmt.SelectExprs[0])) + switch expr { + case "1": + return c.rows([]string{"1"}, [][]interface{}{{int64(1)}}) + case "version()": + return c.rows([]string{"version"}, [][]interface{}{{"milvus-sql-proxy"}}) + case "database()", "current_database()": + return c.rows([]string{"database"}, [][]interface{}{{c.db}}) + } } - - tableName := sqlparser.String(froms[0]) - - tableSchema, err := c.GetCollectinSchema(tableName) + plan, err := planSelect(stmt) + if err != nil { + return err + } + schema, err := c.GetCollectinSchema(plan.table) if err != nil { - return c.writeError(err) + return err + } + schemaFields := map[string]*entity.Field{} + var pk string + for _, f := range schema.Fields { + schemaFields[f.Name] = f + if f.PrimaryKey { + pk = f.Name + } } - // TODO: use real schema - var outputFields []string - var outputFieldsOrder = make(map[string]int) + names := []string{} + types := []entity.FieldType{} if len(stmt.SelectExprs) == 1 && sqlparser.String(stmt.SelectExprs[0]) == "*" { - outputFields = make([]string, len(tableSchema.Fields)) - for i, field := range tableSchema.Fields { - outputFields[i] = field.Name + for _, f := range schema.Fields { + names = append(names, f.Name) + types = append(types, f.DataType) } } else { - outputFields = make([]string, len(stmt.SelectExprs)) - for i, expr := range stmt.SelectExprs { - outputFields[i] = sqlparser.String(expr) + for _, expr := range stmt.SelectExprs { + name := sqlparser.String(expr) + if name != "count(*)" { + name = expr.(*sqlparser.AliasedExpr).Expr.(*sqlparser.ColName).Name.String() + } + switch { + case name == "count(*)": + if plan.vectorField != "" { + return fmt.Errorf("count with vector search is unsupported") + } + names = append(names, name) + types = append(types, entity.FieldTypeInt64) + case name == "_distance" && plan.vectorField != "": + names = append(names, name) + types = append(types, entity.FieldTypeFloat) + default: + f := schemaFields[name] + if f == nil { + return fmt.Errorf("unknown field %s", name) + } + names = append(names, name) + types = append(types, f.DataType) + } } } - - for i, field := range outputFields { - outputFieldsOrder[field] = i + fields := make([]*mysql.Field, len(names)) + for i, n := range names { + fields[i] = resultField(n, types[i]) } - - golog.Info("conn", "handleSelect", "upstream.Query", c.connectionId) - // TODO: other expr, limits - resp, err := c.upstream.Query(c.ctx, tableName, []string{}, "", outputFields, client.WithLimit(100)) + result := func(rows [][]interface{}) error { + r, err := c.buildResultset(fields, names, rows) + if err != nil { + return err + } + r.Fields = fields + return c.writeResultset(c.status, r) + } + if c.describe || plan.limit == 0 { + return result(nil) + } + outputs := []string{} + seen := map[string]bool{} + for _, n := range names { + if n != "_distance" && !seen[n] { + outputs = append(outputs, n) + seen[n] = true + } + } + var cols client.ResultSet + if plan.vectorField != "" { + f := schemaFields[plan.vectorField] + if f == nil || f.DataType != entity.FieldTypeFloatVector { + return fmt.Errorf("search requires a FLOAT_VECTOR field") + } + v, e := fieldValue(f, plan.vectorExpr) + if e != nil { + return e + } + indexes, e := c.upstream.DescribeIndex(c.ctx, plan.table, plan.vectorField) + if e != nil { + return e + } + if len(indexes) == 0 { + return fmt.Errorf("vector field has no index") + } + params := indexes[0].Params() + metric := entity.MetricType(params["metric_type"]) + if metric == "" { + return fmt.Errorf("vector index has no metric_type") + } + sp := &querySearchParams{values: map[string]interface{}{}} + switch strings.ToUpper(string(indexes[0].IndexType())) { + case "HNSW": + ef := plan.limit + plan.offset + if ef < 64 { + ef = 64 + } + sp.values["ef"] = ef + case "IVF_FLAT", "IVF_SQ8", "IVF_PQ": + sp.values["nprobe"] = 16 + } + rs, e := c.upstream.Search(c.ctx, plan.table, nil, plan.filter, outputs, []entity.Vector{entity.FloatVector(v.([]float32))}, plan.vectorField, metric, int(plan.limit), sp, client.WithOffset(plan.offset), client.WithSearchQueryConsistencyLevel(entity.ClStrong)) + if e != nil { + return e + } + if len(rs) == 0 { + return result(nil) + } + r := rs[0] + if r.Err != nil { + return r.Err + } + cols = r.Fields + if r.IDs != nil && cols.GetColumn(pk) == nil { + cols = append(cols, r.IDs) + } + rows := make([][]interface{}, r.ResultCount) + for i := range rows { + rows[i] = make([]interface{}, len(names)) + for j, n := range names { + if n == "_distance" { + if i >= len(r.Scores) { + return fmt.Errorf("incomplete search scores") + } + rows[i][j] = r.Scores[i] + continue + } + col := cols.GetColumn(n) + if col == nil && n == pk { + col = r.IDs + } + if col == nil { + return fmt.Errorf("missing result field %s", n) + } + rows[i][j], err = col.Get(i) + if err != nil { + return err + } + } + } + return result(rows) + } + opts := []client.SearchQueryOptionFunc{client.WithSearchQueryConsistencyLevel(entity.ClStrong)} + if len(names) != 1 || names[0] != "count(*)" { + opts = append(opts, client.WithLimit(plan.limit), client.WithOffset(plan.offset)) + } else if stmt.Limit != nil { + return fmt.Errorf("LIMIT on count(*) is unsupported") + } + cols, err = c.upstream.Query(c.ctx, plan.table, nil, plan.filter, outputs, opts...) if err != nil { - return c.writeError(err) + return err } - golog.Debug("conn", "handleSelect", "upstream.Query finished", c.connectionId, "rows", resp[0].Len(), "columns", len(resp)) - if len(resp) == 0 { - return c.writeOK(nil) + if len(cols) == 0 { + return result(nil) } - - ret := make([][]interface{}, resp[0].Len()) - for i := range ret { - ret[i] = make([]interface{}, len(outputFields)) - for _, column := range resp { - columnIdx, found := outputFieldsOrder[column.Name()] - if !found { - // id will always be returned, ignore - continue + rows := make([][]interface{}, cols[0].Len()) + for i := range rows { + rows[i] = make([]interface{}, len(names)) + for j, n := range names { + col := cols.GetColumn(n) + if col == nil { + return fmt.Errorf("missing result field %s", n) } - ret[i][columnIdx], err = column.Get(i) + rows[i][j], err = col.Get(i) if err != nil { - return c.writeError(err) + return err } } } - golog.Debug("conn", "buildResultset", "buildResultset", c.connectionId, "ret", ret) - r, err := c.buildResultset(nil, outputFields, ret) - if err != nil { - return mysql.NewError(mysql.ER_UNKNOWN_ERROR, errors.Wrap(err, "build resultset failed").Error()) + return result(rows) +} +func resultField(name string, t entity.FieldType) *mysql.Field { + f := &mysql.Field{Name: []byte(name), Charset: 63} + switch t { + case entity.FieldTypeBool: + f.Type = mysql.MYSQL_TYPE_TINY + case entity.FieldTypeInt8, entity.FieldTypeInt16, entity.FieldTypeInt32, entity.FieldTypeInt64: + f.Type = mysql.MYSQL_TYPE_LONGLONG + case entity.FieldTypeFloat, entity.FieldTypeDouble: + f.Type = mysql.MYSQL_TYPE_DOUBLE + default: + f.Type = mysql.MYSQL_TYPE_VAR_STRING + f.Charset = 33 } - golog.Debug("conn", "handleSelect", "writeResultset", c.connectionId, "r", r) - return c.writeResultset(c.status, r) + return f } + +type querySearchParams struct{ values map[string]interface{} } + +func (p *querySearchParams) Params() map[string]interface{} { return p.values } +func (p *querySearchParams) AddRadius(v float64) { p.values["radius"] = v } +func (p *querySearchParams) AddRangeFilter(v float64) { p.values["range_filter"] = v } diff --git a/pkg/select_plan.go b/pkg/select_plan.go new file mode 100644 index 0000000..c481f78 --- /dev/null +++ b/pkg/select_plan.go @@ -0,0 +1,246 @@ +package pkg + +import ( + "fmt" + "github.com/xwb1989/sqlparser" + "regexp" + "strconv" + "strings" +) + +type selectPlan struct { + table, filter string + vectorField string + vectorExpr sqlparser.Expr + limit, offset int64 +} + +var fieldIdentifier = regexp.MustCompile(`^[A-Za-z_][A-Za-z0-9_]*$`) + +// planSelect rejects syntax whose semantics cannot be preserved by scalar Query. +func planSelect(s *sqlparser.Select) (selectPlan, error) { + p := selectPlan{limit: 100} + if s.Distinct != "" || len(s.GroupBy) > 0 || s.Having != nil || len(s.OrderBy) > 0 || s.Lock != "" { + return p, fmt.Errorf("DISTINCT, GROUP BY, HAVING, ORDER BY and locking are not supported") + } + if len(s.From) != 1 { + return p, fmt.Errorf("SELECT requires one collection") + } + a, ok := s.From[0].(*sqlparser.AliasedTableExpr) + if !ok { + return p, fmt.Errorf("joins are not supported") + } + t, ok := a.Expr.(sqlparser.TableName) + if !ok || !a.As.IsEmpty() || a.Hints != nil || !t.Qualifier.IsEmpty() { + return p, fmt.Errorf("subqueries, table aliases, hints and qualified tables are not supported") + } + p.table = t.Name.String() + for _, e := range s.SelectExprs { + if star, ok := e.(*sqlparser.StarExpr); ok { + if len(s.SelectExprs) != 1 || !star.TableName.IsEmpty() { + return p, fmt.Errorf("only an unqualified standalone * is supported") + } + continue + } + a, ok := e.(*sqlparser.AliasedExpr) + if !ok || !a.As.IsEmpty() { + return p, fmt.Errorf("projection aliases are not supported") + } + if sqlparser.String(a.Expr) == "count(*)" && len(s.SelectExprs) == 1 { + continue + } + c, ok := a.Expr.(*sqlparser.ColName) + if !ok || !c.Qualifier.IsEmpty() { + return p, fmt.Errorf("projection must be an unqualified field") + } + } + var err error + if s.Where != nil { + expr, field, vector, e := extractVector(s.Where.Expr) + if e != nil { + return p, e + } + p.vectorField = field + p.vectorExpr = vector + if expr != nil { + p.filter, err = scalarFilter(expr) + } + if err != nil { + return p, err + } + } + if s.Limit != nil { + p.limit, err = selectInteger(s.Limit.Rowcount) + if err != nil { + return p, err + } + if s.Limit.Offset != nil { + p.offset, err = selectInteger(s.Limit.Offset) + if err != nil { + return p, err + } + } + } + if p.limit > 16384 || p.offset > 16384 || p.limit+p.offset > 16384 { + return p, fmt.Errorf("LIMIT + OFFSET must not exceed 16384") + } + return p, nil +} +func selectInteger(e sqlparser.Expr) (int64, error) { + v, ok := e.(*sqlparser.SQLVal) + if !ok || v.Type != sqlparser.IntVal { + return 0, fmt.Errorf("LIMIT and OFFSET must be non-negative integer literals") + } + n, err := strconv.ParseInt(string(v.Val), 10, 64) + if err != nil || n < 0 { + return 0, fmt.Errorf("invalid LIMIT or OFFSET") + } + return n, nil +} +func scalarFilter(e sqlparser.Expr) (string, error) { + switch v := e.(type) { + case *sqlparser.AndExpr: + return filterPair(v.Left, v.Right, "and") + case *sqlparser.OrExpr: + return filterPair(v.Left, v.Right, "or") + case *sqlparser.NotExpr: + x, err := scalarFilter(v.Expr) + return "not (" + x + ")", err + case *sqlparser.ParenExpr: + x, err := scalarFilter(v.Expr) + return "(" + x + ")", err + case *sqlparser.ComparisonExpr: + op := v.Operator + switch op { + case "=": + op = "==" + case "!=", "<>": + op = "!=" + case "in", "not in": + l, err := filterValue(v.Left) + if err != nil { + return "", err + } + tuple, ok := v.Right.(sqlparser.ValTuple) + if !ok { + return "", fmt.Errorf("IN requires a literal list") + } + items := []string{} + for _, x := range tuple { + switch x.(type) { + case *sqlparser.SQLVal, sqlparser.BoolVal, *sqlparser.UnaryExpr: + default: + return "", fmt.Errorf("IN requires literals") + } + item, err := filterValue(x) + if err != nil { + return "", err + } + items = append(items, item) + } + return l + " " + op + " [" + strings.Join(items, ",") + "]", nil + case "like": + if _, ok := v.Right.(*sqlparser.SQLVal); !ok { + return "", fmt.Errorf("LIKE requires a string literal") + } + case "<", "<=", ">", ">=": + default: + return "", fmt.Errorf("unsupported WHERE operator %s", op) + } + if v.Escape != nil { + return "", fmt.Errorf("ESCAPE is not supported") + } + l, err := filterValue(v.Left) + if err != nil { + return "", err + } + r, err := filterValue(v.Right) + if err != nil { + return "", err + } + return l + " " + op + " " + r, nil + default: + return "", fmt.Errorf("unsupported WHERE expression %T", e) + } +} +func filterPair(l, r sqlparser.Expr, op string) (string, error) { + a, err := scalarFilter(l) + if err != nil { + return "", err + } + b, err := scalarFilter(r) + if err != nil { + return "", err + } + return "(" + a + " " + op + " " + b + ")", nil +} +func filterValue(e sqlparser.Expr) (string, error) { + switch v := e.(type) { + case *sqlparser.ColName: + if v.Qualifier.IsEmpty() && fieldIdentifier.MatchString(v.Name.String()) { + return v.Name.String(), nil + } + case *sqlparser.SQLVal: + switch v.Type { + case sqlparser.StrVal: + return strconv.Quote(string(v.Val)), nil + case sqlparser.IntVal, sqlparser.FloatVal: + return string(v.Val), nil + } + case sqlparser.BoolVal: + return strconv.FormatBool(bool(v)), nil + case *sqlparser.UnaryExpr: + if v.Operator == "-" || v.Operator == "+" { + if n, ok := v.Expr.(*sqlparser.SQLVal); ok && (n.Type == sqlparser.IntVal || n.Type == sqlparser.FloatVal) { + return v.Operator + string(n.Val), nil + } + } + } + return "", fmt.Errorf("unsupported WHERE value %T", e) +} + +// Vector predicates are allowed only as positive AND conjuncts. Treating ANN as +// a scalar predicate under OR/NOT would change the meaning of the SQL query. +func extractVector(e sqlparser.Expr) (sqlparser.Expr, string, sqlparser.Expr, error) { + switch v := e.(type) { + case *sqlparser.ParenExpr: + x, f, q, err := extractVector(v.Expr) + if err != nil || x == nil { + return x, f, q, err + } + return &sqlparser.ParenExpr{Expr: x}, f, q, nil + case *sqlparser.AndExpr: + l, lf, lq, err := extractVector(v.Left) + if err != nil { + return nil, "", nil, err + } + r, rf, rq, err := extractVector(v.Right) + if err != nil { + return nil, "", nil, err + } + if lf != "" && rf != "" { + return nil, "", nil, fmt.Errorf("only one vector predicate is supported") + } + if rf != "" { + lf, lq = rf, rq + } + if l == nil { + return r, lf, lq, nil + } + if r == nil { + return l, lf, lq, nil + } + return &sqlparser.AndExpr{Left: l, Right: r}, lf, lq, nil + case *sqlparser.ComparisonExpr: + if v.Operator == "like" { + if fn, ok := v.Right.(*sqlparser.FuncExpr); ok && strings.EqualFold(fn.Name.String(), "json_vector") { + col, ok := v.Left.(*sqlparser.ColName) + if !ok || !col.Qualifier.IsEmpty() || v.Escape != nil { + return nil, "", nil, fmt.Errorf("invalid vector predicate") + } + return nil, col.Name.String(), v.Right, nil + } + } + } + return e, "", nil, nil +} diff --git a/pkg/select_test.go b/pkg/select_test.go new file mode 100644 index 0000000..72f714a --- /dev/null +++ b/pkg/select_test.go @@ -0,0 +1,146 @@ +package pkg + +import ( + "context" + "fmt" + + "github.com/milvus-io/milvus-sdk-go/v2/client" + "github.com/milvus-io/milvus-sdk-go/v2/entity" + "github.com/xwb1989/sqlparser" + + "reflect" + "testing" +) + +func parsedSelect(t *testing.T, q string) *sqlparser.Select { + t.Helper() + s, e := sqlparser.Parse(q) + if e != nil { + t.Fatal(e) + } + return s.(*sqlparser.Select) +} +func TestSelectPlan(t *testing.T) { + for _, tc := range []struct { + sql, filter string + limit, offset int64 + }{ + {"select * from docs", "", 100, 0}, + {"select id from docs where id=1 limit 3 offset 2", "id == 1", 3, 2}, + {"select id from docs where name='a' and (id<>2 or id>=3) limit 2,4", `(name == "a" and ((id != 2 or id >= 3)))`, 4, 2}, + {"select id from docs where not (id < -1)", "not ((id < -1))", 100, 0}, + {"select id from docs where enabled=true", "enabled == true", 100, 0}, + {"select id from docs limit 0", "", 0, 0}, + } { + t.Run(tc.sql, func(t *testing.T) { + p, e := planSelect(parsedSelect(t, tc.sql)) + if e != nil { + t.Fatal(e) + } + if p.table != "docs" || p.filter != tc.filter || p.limit != tc.limit || p.offset != tc.offset { + t.Fatalf("unexpected plan: %+v", p) + } + }) + } +} +func TestUnsupportedSelect(t *testing.T) { + for _, q := range []string{ + "select id from docs order by id", "select distinct id from docs", "select id from docs group by id", "select id from docs having id=1", "select id from docs for update", + "select * from docs, other", "select * from docs join other on docs.id=other.id", "select * from (select * from docs) d", "select * from docs d", "select * from db.docs", "select docs.* from docs", "select *, id from docs", "select id as x from docs", "select docs.id from docs", "select * from docs use index (idx)", + "select id from docs where id is null", "select id from docs where id=abs(1)", "select id from docs where abs(id)=1", "select id from docs where id=1 and id is null", "select id from docs where id is null or id=1", "select id from docs where db.id=1", "select id from docs limit :n", "select id from docs limit 9999999999999999999999999", "select id from docs limit 1 offset :n", + } { + t.Run(q, func(t *testing.T) { + if _, e := planSelect(parsedSelect(t, q)); e == nil { + t.Fatal("expected rejection") + } + }) + } +} + +type selectMock struct { + queryErr, schemaErr error + client.Client + filter string + fields []string + opts client.SearchQueryOption + calls int + result client.ResultSet +} + +func (m *selectMock) DescribeCollection(context.Context, string) (*entity.Collection, error) { + if m.schemaErr != nil { + return nil, m.schemaErr + } + return &entity.Collection{Schema: &entity.Schema{Fields: []*entity.Field{{Name: "id"}, {Name: "name"}}}}, nil +} +func (m *selectMock) Query(_ context.Context, _ string, _ []string, filter string, fields []string, opts ...client.SearchQueryOptionFunc) (client.ResultSet, error) { + m.calls++ + m.filter = filter + m.fields = fields + for _, o := range opts { + o(&m.opts) + } + return m.result, m.queryErr +} + +func TestSelectSDKAndWire(t *testing.T) { + for _, empty := range []bool{false, true} { + m := &selectMock{} + if !empty { + m.result = client.ResultSet{entity.NewColumnInt64("id", []int64{7}), entity.NewColumnVarChar("name", []string{"alice"})} + } + + c := &ClientConn{ctx: context.Background(), upstream: m} + if e := c.handleSelect(parsedSelect(t, "select name,id from docs where id=7 limit 2 offset 3"), nil); e != nil { + t.Fatal(e) + } + if m.filter != "id == 7" || m.opts.Limit != 2 || m.opts.Offset != 3 || !reflect.DeepEqual(m.fields, []string{"name", "id"}) { + t.Fatalf("wrong SDK call: %+v", m) + } + if c.result == nil || len(c.result.Fields) != 2 { + t.Fatal("expected 2-column result") + } + if !empty && c.result.Values[0][0] != "alice" { + t.Fatal("missing result row") + } + } +} +func TestSelectZeroLimit(t *testing.T) { + m := &selectMock{} + + c := &ClientConn{ctx: context.Background(), upstream: m} + if e := c.handleSelect(parsedSelect(t, "select * from docs limit 0"), nil); e != nil { + t.Fatal(e) + } + if m.calls != 0 { + t.Fatal("LIMIT 0 queried upstream") + } +} + +func TestSelectErrors(t *testing.T) { + for _, tc := range []struct { + q string + m *selectMock + }{ + {"select id from docs order by id", &selectMock{}}, + {"select id from docs", &selectMock{schemaErr: fmt.Errorf("schema unavailable")}}, + {"select id from docs", &selectMock{queryErr: fmt.Errorf("query unavailable")}}, + } { + + c := &ClientConn{ctx: context.Background(), upstream: tc.m} + if e := c.handleSelect(parsedSelect(t, tc.q), nil); e == nil { + t.Fatal("expected error") + } + } +} +func TestSelectRepeatedProjection(t *testing.T) { + m := &selectMock{result: client.ResultSet{entity.NewColumnInt64("id", []int64{7})}} + + c := &ClientConn{ctx: context.Background(), upstream: m} + if e := c.handleSelect(parsedSelect(t, "select id,id from docs"), nil); e != nil { + t.Fatal(e) + } + if !reflect.DeepEqual(c.result.Values[0], []interface{}{int64(7), int64(7)}) { + t.Fatal("repeated projection lost a value") + } +} diff --git a/pkg/server.go b/pkg/server.go index aebe4fe..1a997eb 100644 --- a/pkg/server.go +++ b/pkg/server.go @@ -1,299 +1,140 @@ -// partially copied & changed from : https://github.com/flike/kingshard/blob/master/proxy/server/server.go - -// Copyright 2016 The kingshard Authors. All rights reserved. -// -// Licensed under the Apache License, Version 2.0 (the "License"): you may -// not use this file except in compliance with the License. You may obtain -// a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, software -// distributed under the License is distributed on an "AS IS" BASIS, WITHOUT -// WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. See the -// License for the specific language governing permissions and limitations -// under the License. package pkg import ( "context" + "crypto/tls" + "errors" "fmt" "net" - "runtime" - "strconv" - "sync/atomic" + "runtime/debug" + "sync" "time" - "github.com/flike/kingshard/mysql" "github.com/milvus-io/milvus-sdk-go/v2/client" - - "sync" - - "github.com/flike/kingshard/backend" - "github.com/flike/kingshard/core/errors" - "github.com/flike/kingshard/core/golog" - "github.com/flike/kingshard/proxy/router" - pkgErr "github.com/pkg/errors" -) - -type Schema struct { - nodes map[string]*backend.Node - rule *router.Router -} - -type BlacklistSqls struct { - sqls map[string]string - sqlsLen int -} - -const ( - Offline = iota - Online - Unknown ) type Server struct { - cfg *Config - addr string - users map[string]string //user : psw - - statusIndex int32 - status [2]int32 - logSqlIndex int32 - logSql [2]string - slowLogTimeIndex int32 - slowLogTime [2]int - - counter *Counter - nodes map[string]*backend.Node - - listener net.Listener - running bool - - configUpdateMutex sync.RWMutex - configVer uint32 -} - -func (s *Server) Status() string { - var status string - switch s.status[s.statusIndex] { - case Online: - status = "online" - case Offline: - status = "offline" - case Unknown: - status = "unknown" - default: - status = "unknown" - } - return status + cfg *Config + listeners []net.Listener + modes []string + mu sync.Mutex + conns map[net.Conn]context.CancelFunc + closed bool + wg sync.WaitGroup + tlsConfig *tls.Config + newSession func(context.Context) (*ClientConn, error) } func NewServer(cfg *Config) (*Server, error) { - s := new(Server) - - s.cfg = cfg - s.counter = new(Counter) - s.addr = cfg.Addr - - golog.Info("server", "NewServer", "addr", 0, "addr", s.addr) - atomic.StoreInt32(&s.statusIndex, 0) - s.status[s.statusIndex] = Online - s.configVer = 0 - - var err error - netProto := "tcp" - - s.listener, err = net.Listen(netProto, s.addr) - - if err != nil { + if err := cfg.defaults(); err != nil { return nil, err } - - golog.Info("server", "NewServer", "Server running", 0, - "netProto", - netProto, - "address", - s.addr) - return s, nil -} - -func (s *Server) flushCounter() { - for { - s.counter.FlushCounter() - time.Sleep(1 * time.Second) + s := &Server{cfg: cfg, conns: make(map[net.Conn]context.CancelFunc)} + if cfg.TLSCert != "" { + cert, err := tls.LoadX509KeyPair(cfg.TLSCert, cfg.TLSKey) + if err != nil { + return nil, err + } + s.tlsConfig = &tls.Config{Certificates: []tls.Certificate{cert}, MinVersion: tls.VersionTLS12} } -} - -func (s *Server) newClientConn(ctx context.Context, co net.Conn) (*ClientConn, error) { - c := new(ClientConn) - c.ctx = ctx - tcpConn := co.(*net.TCPConn) - - //SetNoDelay controls whether the operating system should delay packet transmission - // in hopes of sending fewer packets (Nagle's algorithm). - // The default is true (no delay), - // meaning that data is sent as soon as possible after a Write. - //I set this option false. - tcpConn.SetNoDelay(false) - c.c = tcpConn - - c.pkg = mysql.NewPacketIO(tcpConn) - // c.proxy = s - - var err error - cfg := s.cfg - milvusCfg := client.Config{ - Address: cfg.Milvus.Address, - Username: cfg.Milvus.Username, - Password: cfg.Milvus.Password, - APIKey: cfg.Milvus.APIKey, + s.newSession = func(ctx context.Context) (*ClientConn, error) { + upstream, err := client.NewClient(ctx, client.Config{Address: cfg.Milvus.Address, Username: cfg.Milvus.Username, Password: cfg.Milvus.Password, APIKey: cfg.Milvus.APIKey, EnableTLSAuth: cfg.Milvus.EnableTLSAuth}) + if err != nil { + return nil, err + } + return NewSession(ctx, upstream), nil } - c.upstream, err = client.NewClient(ctx, milvusCfg) - if err != nil { - return nil, pkgErr.Wrap(err, "connect to milvus failed") + bind := func(mode, addr string) error { + l, err := net.Listen("tcp", addr) + if err != nil { + return err + } + s.listeners = append(s.listeners, l) + s.modes = append(s.modes, mode) + return nil } - - c.pkg.Sequence = 0 - - c.connectionId = atomic.AddUint32(&baseConnId, 1) - - c.status = mysql.SERVER_STATUS_AUTOCOMMIT - - c.salt, _ = mysql.RandomBuf(20) - - c.closed = false - - c.charset = mysql.DEFAULT_CHARSET - c.collation = mysql.DEFAULT_COLLATION_ID - - c.stmtId = 0 - c.stmts = make(map[uint32]*Stmt) - - return c, nil -} - -func (s *Server) onConn(c net.Conn) { - ctx, cancel := context.WithCancel(context.Background()) - defer cancel() - s.counter.IncrClientConns() - conn, err := s.newClientConn(ctx, c) //新建一个conn - if err != nil { - conn.writeError(err) - conn.Close() - return + if cfg.Mode == "mysql" || cfg.Mode == "both" { + if err := bind("mysql", cfg.Addr); err != nil { + s.Close() + return nil, err + } } - - defer func() { - err := recover() - if err != nil { - const size = 4096 - buf := make([]byte, size) - buf = buf[:runtime.Stack(buf, false)] //获得当前goroutine的stacktrace - golog.Error("server", "onConn", "error", 0, - "remoteAddr", c.RemoteAddr().String(), - "stack", string(buf), - ) + if cfg.Mode == "postgres" || cfg.Mode == "both" { + if err := bind("postgres", cfg.PostgresAddr); err != nil { + s.Close() + return nil, err } - - conn.Close() - s.counter.DecrClientConns() - }() - - if err := conn.Handshake(); err != nil { - golog.Error("server", "onConn", err.Error(), 0) - conn.writeError(err) - conn.Close() - return } - - conn.Run() + return s, nil } - -func (s *Server) ChangeProxy(v string) error { - var status int32 - switch v { - case "online": - status = Online - case "offline": - status = Offline - default: - status = Unknown - } - if status == Unknown { - return errors.ErrCmdUnsupport +func (s *Server) Run() error { + errs := make(chan error, len(s.listeners)) + for i, l := range s.listeners { + mode := s.modes[i] + go func() { errs <- s.serve(l, mode) }() } - - if s.statusIndex == 0 { - s.status[1] = status - atomic.StoreInt32(&s.statusIndex, 1) - } else { - s.status[0] = status - atomic.StoreInt32(&s.statusIndex, 0) + var first error + for range s.listeners { + if err := <-errs; err != nil && first == nil { + first = err + s.Close() + } } - - return nil + s.wg.Wait() + return first } - -func (s *Server) Run() error { - s.running = true - - // flush counter - go s.flushCounter() - - for s.running { - conn, err := s.listener.Accept() +func (s *Server) serve(l net.Listener, mode string) error { + for { + co, err := l.Accept() if err != nil { - golog.Error("server", "Run", err.Error(), 0) - continue + if errors.Is(err, net.ErrClosed) { + return nil + } + return err } - - go s.onConn(conn) + s.mu.Lock() + if s.closed { + s.mu.Unlock() + co.Close() + return nil + } + ctx, cancel := context.WithCancel(context.Background()) + s.conns[co] = cancel + s.wg.Add(1) + s.mu.Unlock() + go func() { + defer s.wg.Done() + defer co.Close() + defer func() { s.mu.Lock(); delete(s.conns, co); s.mu.Unlock() }() + // Malformed clients may cause third-party protocol decoders to panic. Keep + // a single connection from taking down the listener. + defer func() { + if r := recover(); r != nil { + fmt.Printf("%s connection failed: %v\n%s", mode, r, debug.Stack()) + } + }() + defer cancel() + co.SetDeadline(time.Now().Add(15 * time.Second)) + if mode == "postgres" { + s.servePostgres(ctx, co) + } else { + s.serveMySQL(ctx, co) + } + }() } - - return nil } - func (s *Server) Close() { - s.running = false - if s.listener != nil { - s.listener.Close() + s.mu.Lock() + defer s.mu.Unlock() + if s.closed { + return } -} - -func (s *Server) GetMonitorData() map[string]map[string]string { - data := make(map[string]map[string]string) - - // get all node's monitor data - for _, node := range s.nodes { - //get master monitor data - dbData := make(map[string]string) - idleConns, cacheConns, pushConnCount, popConnCount := node.Master.ConnCount() - - dbData["idleConn"] = strconv.Itoa(idleConns) - dbData["cacheConns"] = strconv.Itoa(cacheConns) - dbData["pushConnCount"] = strconv.FormatInt(pushConnCount, 10) - dbData["popConnCount"] = strconv.FormatInt(popConnCount, 10) - dbData["maxConn"] = fmt.Sprintf("%d", node.Cfg.MaxConnNum) - dbData["type"] = "master" - - data[node.Master.Addr()] = dbData - - //get all slave monitor data - for _, slaveNode := range node.Slave { - slaveDbData := make(map[string]string) - idleConns, cacheConns, pushConnCount, popConnCount := slaveNode.ConnCount() - - slaveDbData["idleConn"] = strconv.Itoa(idleConns) - slaveDbData["cacheConns"] = strconv.Itoa(cacheConns) - slaveDbData["pushConnCount"] = strconv.FormatInt(pushConnCount, 10) - slaveDbData["popConnCount"] = strconv.FormatInt(popConnCount, 10) - slaveDbData["maxConn"] = fmt.Sprintf("%d", node.Cfg.MaxConnNum) - slaveDbData["type"] = "slave" - - data[slaveNode.Addr()] = slaveDbData - } + s.closed = true + for _, l := range s.listeners { + l.Close() + } + for co, cancel := range s.conns { + cancel() + co.Close() } - - return data } diff --git a/pkg/show.go b/pkg/show.go index d992b74..c67b503 100644 --- a/pkg/show.go +++ b/pkg/show.go @@ -18,12 +18,15 @@ const ( func (c *ClientConn) handleShow(stmt *sqlparser.Show, args []interface{}) error { switch strings.ToUpper(stmt.Type) { case DatabasesStr: + if c.describe { + return c.rows([]string{DatabasesStr}, nil) + } ret, err := c.upstream.ListDatabases(c.ctx) if err != nil { return mysql.NewError(mysql.ER_ABORTING_CONNECTION, errors.Wrap(err, "list databases failed").Error()) } if len(ret) == 0 { - return c.writeOK(nil) + return c.rows([]string{DatabasesStr}, nil) } values := databasesToValues(ret) r, err := c.buildResultset(nil, []string{DatabasesStr}, values) @@ -32,6 +35,9 @@ func (c *ClientConn) handleShow(stmt *sqlparser.Show, args []interface{}) error } return c.writeResultset(c.status, r) case TableStr: + if c.describe { + return c.rows([]string{TableStr}, nil) + } ret, err := c.upstream.ListCollections(c.ctx) if err != nil { return mysql.NewError(mysql.ER_ABORTING_CONNECTION, errors.Wrap(err, "list collections failed").Error()) diff --git a/pkg/types.go b/pkg/types.go index 5af8bad..60e0d14 100644 --- a/pkg/types.go +++ b/pkg/types.go @@ -1,8 +1,9 @@ package pkg -import "github.com/flike/kingshard/proxy/server" +import "github.com/milvus-io/milvus-sdk-go/v2/entity" -type ( - Counter = server.Counter - Stmt = server.Stmt -) +// MilvusSchema includes collection fields and its shard count. +type MilvusSchema struct { + *entity.Schema + ShardNum int32 +} diff --git a/testdata/embedEtcd.yaml b/testdata/embedEtcd.yaml new file mode 100644 index 0000000..32954fa --- /dev/null +++ b/testdata/embedEtcd.yaml @@ -0,0 +1,5 @@ +listen-client-urls: http://0.0.0.0:2379 +advertise-client-urls: http://0.0.0.0:2379 +quota-backend-bytes: 4294967296 +auto-compaction-mode: revision +auto-compaction-retention: '1000' From 2138f88a74468c9318386900ed397f58aa997238 Mon Sep 17 00:00:00 2001 From: "shaoyue.chen" Date: Fri, 18 Sep 2026 20:17:55 +0800 Subject: [PATCH 2/5] Validate lifecycle routing and fix MySQL shutdown and integration decoding --- pkg/conn_resultvalue.go | 71 +++++------------- pkg/integration_test.go | 8 ++ pkg/lifecycle_test.go | 162 ++++++++++++++++++++++++++++++++++++++++ pkg/mysql.go | 6 +- 4 files changed, 192 insertions(+), 55 deletions(-) create mode 100644 pkg/lifecycle_test.go diff --git a/pkg/conn_resultvalue.go b/pkg/conn_resultvalue.go index 4c82853..70bdf9c 100644 --- a/pkg/conn_resultvalue.go +++ b/pkg/conn_resultvalue.go @@ -95,69 +95,28 @@ func formatField(field *mysql.Field, value interface{}) error { return nil } +// Result construction is transport-neutral: adapters encode rows on the wire. func (c *ClientConn) buildResultset(fields []*mysql.Field, names []string, values [][]interface{}) (*mysql.Resultset, error) { - var ExistFields bool - r := new(mysql.Resultset) - - r.Fields = make([]*mysql.Field, len(names)) - r.FieldNames = make(map[string]int, len(names)) - - //use the field def that get from true database - if len(fields) != 0 { - if len(r.Fields) == len(fields) { - ExistFields = true - } else { - return nil, errors.ErrInvalidArgument - } + if len(fields) != 0 && len(fields) != len(names) { + return nil, errors.ErrInvalidArgument } - - if len(values) == 0 { - return newEmptyResultset(names), nil + r := newEmptyResultset(names) + if len(fields) > 0 { + r.Fields = fields } - - var b []byte - var err error - - for i, vs := range values { - if len(vs) != len(r.Fields) { - return nil, fmt.Errorf("row %d has %d column not equal %d", i, len(vs), len(r.Fields)) + for i, row := range values { + if len(row) != len(names) { + return nil, fmt.Errorf("row %d has %d columns, expected %d", i, len(row), len(names)) } - - var row []byte - for j, value := range vs { - //列的定义 - if i == 0 { - if ExistFields { - r.Fields[j] = fields[j] - r.FieldNames[string(r.Fields[j].Name)] = j - } else { - field := &mysql.Field{} - r.Fields[j] = field - field.Name = hack.Slice(names[j]) - r.FieldNames[names[j]] = j - if err = formatField(field, value); err != nil { - return nil, err - } + if i == 0 && len(fields) == 0 { + for j, v := range row { + if err := formatField(r.Fields[j], v); err != nil { + return nil, err } - - } - if value == nil { - row = append(row, 0xfb) - continue - } - b, err = formatValue(value) - if err != nil { - return nil, err } - - row = append(row, mysql.PutLengthEncodedString(b)...) } - - r.RowDatas = append(r.RowDatas, row) } - //assign the values to the result r.Values = values - return r, nil } @@ -169,9 +128,13 @@ func (c *ClientConn) writeResultset(status uint16, r *mysql.Resultset) error { func newEmptyResultset(fields []string) *mysql.Resultset { r := new(mysql.Resultset) r.Fields = make([]*mysql.Field, len(fields)) + r.FieldNames = make(map[string]int, len(fields)) for i := range fields { r.Fields[i] = &mysql.Field{} r.Fields[i].Name = hack.Slice(fields[i]) + r.Fields[i].Type = mysql.MYSQL_TYPE_VAR_STRING + r.Fields[i].Charset = 33 + r.FieldNames[fields[i]] = i } r.Values = make([][]interface{}, 0) diff --git a/pkg/integration_test.go b/pkg/integration_test.go index 2bd9b43..8ba377f 100644 --- a/pkg/integration_test.go +++ b/pkg/integration_test.go @@ -77,6 +77,11 @@ func TestMilvusIntegration(t *testing.T) { if e := rs.Scan(ptrs...); e != nil { return nil, e } + for i, v := range values { + if b, ok := v.([]byte); ok { + values[i] = string(b) + } + } out = append(out, values) } return out, rs.Err() @@ -121,6 +126,9 @@ func TestMilvusIntegration(t *testing.T) { } mustExec("CREATE DATABASE " + dbname) defer func() { + exec("USE " + dbname) + exec("RELEASE TABLE items") + exec("DROP TABLE items") exec("USE default") if e := exec("DROP DATABASE " + dbname); e != nil { t.Error(e) diff --git a/pkg/lifecycle_test.go b/pkg/lifecycle_test.go new file mode 100644 index 0000000..0e9783e --- /dev/null +++ b/pkg/lifecycle_test.go @@ -0,0 +1,162 @@ +package pkg + +import ( + "context" + "errors" + "reflect" + "testing" + + "github.com/milvus-io/milvus-sdk-go/v2/client" + "github.com/milvus-io/milvus-sdk-go/v2/entity" +) + +func (m *operationMock) CreateDatabase(_ context.Context, name string, _ ...client.CreateDatabaseOption) error { + m.calls = append(m.calls, "create db "+name) + return m.fail +} +func (m *operationMock) DropDatabase(_ context.Context, name string, _ ...client.DropDatabaseOption) error { + m.calls = append(m.calls, "drop db "+name) + return m.fail +} +func (m *operationMock) ListDatabases(context.Context) ([]entity.Database, error) { + return []entity.Database{{Name: "default"}}, m.fail +} +func (m *operationMock) ListCollections(context.Context) ([]*entity.Collection, error) { + return []*entity.Collection{{Name: "items"}}, m.fail +} +func (m *operationMock) LoadCollection(_ context.Context, name string, async bool, _ ...client.LoadCollectionOption) error { + if async { + return errors.New("expected synchronous load") + } + m.calls = append(m.calls, "load "+name) + return m.fail +} +func (m *operationMock) ReleaseCollection(_ context.Context, name string, _ ...client.ReleaseCollectionOption) error { + m.calls = append(m.calls, "release "+name) + return m.fail +} +func (m *operationMock) Flush(_ context.Context, name string, async bool, _ ...client.FlushOption) error { + if async { + return errors.New("expected synchronous flush") + } + m.calls = append(m.calls, "flush "+name) + return m.fail +} +func (m *operationMock) CreateIndex(_ context.Context, table, field string, idx entity.Index, async bool, _ ...client.IndexOption) error { + if async { + return errors.New("expected synchronous index") + } + m.calls = append(m.calls, "index "+table+" "+field+" "+idx.Name()+" "+string(idx.IndexType())+" "+idx.Params()["metric_type"]) + return m.fail +} +func (m *operationMock) DropIndex(_ context.Context, table, field string, _ ...client.IndexOption) error { + m.calls = append(m.calls, "drop index "+table) + return m.fail +} +func (m *operationMock) DescribeIndex(context.Context, string, string, ...client.IndexOption) ([]entity.Index, error) { + return []entity.Index{entity.NewGenericIndex("vecidx", entity.HNSW, map[string]string{"metric_type": "COSINE"})}, m.fail +} +func (m *operationMock) CreatePartition(_ context.Context, table, partition string, _ ...client.CreatePartitionOption) error { + m.calls = append(m.calls, "create partition "+table+" "+partition) + return m.fail +} +func (m *operationMock) DropPartition(_ context.Context, table, partition string, _ ...client.DropPartitionOption) error { + m.calls = append(m.calls, "drop partition "+table+" "+partition) + return m.fail +} +func (m *operationMock) LoadPartitions(_ context.Context, table string, parts []string, async bool, _ ...client.LoadPartitionsOption) error { + m.calls = append(m.calls, "load partition "+table+" "+parts[0]) + return m.fail +} +func (m *operationMock) ReleasePartitions(_ context.Context, table string, parts []string, _ ...client.ReleasePartitionsOption) error { + m.calls = append(m.calls, "release partition "+table+" "+parts[0]) + return m.fail +} +func (m *operationMock) ShowPartitions(context.Context, string) ([]*entity.Partition, error) { + return []*entity.Partition{{Name: "_default"}}, m.fail +} +func (m *operationMock) Search(_ context.Context, table string, _ []string, filter string, fields []string, vecs []entity.Vector, field string, metric entity.MetricType, k int, params entity.SearchParam, opts ...client.SearchQueryOptionFunc) ([]client.SearchResult, error) { + m.queryFilter = filter + m.calls = append(m.calls, "search") + if table != "items" || field != "vec" || metric != entity.COSINE || len(vecs) != 1 || k != 2 || params.Params()["ef"] != int64(64) { + return nil, errors.New("unexpected search arguments") + } + return []client.SearchResult{{ResultCount: 1, IDs: entity.NewColumnInt64("id", []int64{7}), Scores: []float32{0.75}, Fields: client.ResultSet{entity.NewColumnVarChar("name", []string{"one"})}}}, m.fail +} +func TestLifecycleRouting(t *testing.T) { + ctx := context.Background() + m := &operationMock{schema: testSchema()} + c := NewSession(ctx, m) + cases := []struct{ sql, call string }{ + {"CREATE DATABASE demo", "create db demo"}, {"DROP DATABASE demo", "drop db demo"}, {"LOAD TABLE items", "load items"}, {"RELEASE TABLE items", "release items"}, {"FLUSH TABLE items", "flush items"}, {"CREATE INDEX idx ON items (vec) USING FLAT", "index items vec idx FLAT L2"}, {"CREATE INDEX idx ON items (vec) USING HNSW WITH (metric_type='COSINE', M=16, efConstruction=100)", "index items vec idx HNSW COSINE"}, {"CREATE INDEX idx ON items (vec) USING IVF_FLAT", "index items vec idx IVF_FLAT L2"}, {"CREATE INDEX idx ON items (name) USING INVERTED", "index items name idx INVERTED "}, {"DROP INDEX idx ON items", "drop index items"}, {"CREATE PARTITION p ON items", "create partition items p"}, {"DROP PARTITION p ON items", "drop partition items p"}, {"LOAD PARTITION p ON items", "load partition items p"}, {"RELEASE PARTITION p ON items", "release partition items p"}, + } + for _, tc := range cases { + t.Run(tc.sql, func(t *testing.T) { + m.calls = nil + if _, e := c.Describe(ctx, tc.sql); e != nil { + t.Fatal(e) + } + if len(m.calls) != 0 { + t.Fatal("Describe mutated upstream", m.calls) + } + if _, e := c.Execute(ctx, tc.sql); e != nil { + t.Fatal(e) + } + if !reflect.DeepEqual(m.calls, []string{tc.call}) { + t.Fatal(m.calls) + } + m.fail = errors.New("upstream") + if _, e := c.Execute(ctx, tc.sql); e == nil { + t.Fatal("upstream error lost") + } + m.fail = nil + }) + } + for _, q := range []string{"SHOW DATABASES", "SHOW TABLES", "SHOW PARTITIONS FROM items", "DESCRIBE items"} { + r, e := c.Execute(ctx, q) + if e != nil || r.Resultset == nil || len(r.Values) == 0 { + t.Fatal(q, r, e) + } + d, e := c.Describe(ctx, q) + if e != nil || len(d.Fields) != len(r.Fields) { + t.Fatal(q, d, e) + } + for i := range d.Fields { + if pgOID(d.Fields[i]) != pgOID(r.Fields[i]) { + t.Fatal("metadata mismatch", q) + } + } + m.fail = errors.New("upstream") + if _, e := c.Execute(ctx, q); e == nil { + t.Fatal("upstream error lost") + } + m.fail = nil + } + for _, q := range []string{"CREATE INDEX idx ON items (vec) USING FLAT WITH (metric_type='BOGUS')", "CREATE INDEX idx ON items (vec) USING FLAT WITH (bogus=1)", "CREATE INDEX idx ON items (vec) USING FLAT WITH (nlist)", "CREATE INDEX idx ON items (vec) USING FLAT WITH (nlist=1,nlist=2)"} { + if _, e := c.Execute(ctx, q); e == nil { + t.Fatal(q) + } + } +} +func TestVectorSearch(t *testing.T) { + m := &operationMock{schema: testSchema()} + c := NewSession(context.Background(), m) + r, e := c.Execute(context.Background(), "SELECT id,name,_distance FROM items WHERE vec LIKE json_vector('[1,2,3]') AND enabled=true LIMIT 2") + if e != nil { + t.Fatal(e) + } + if !reflect.DeepEqual(r.Values, [][]interface{}{{int64(7), "one", float32(.75)}}) || m.queryFilter != "enabled == true" { + t.Fatal(r.Values, m.queryFilter) + } + for _, q := range []string{"SELECT id FROM items WHERE vec LIKE json_vector('[1,2,3]') OR id=1 LIMIT 2", "SELECT id FROM items WHERE NOT (vec LIKE json_vector('[1,2,3]')) LIMIT 2", "SELECT id FROM items WHERE vec LIKE json_vector('[1,2,3]') AND vec LIKE json_vector('[1,2,3]') LIMIT 2", "SELECT id FROM items WHERE name LIKE json_vector('[1,2,3]') LIMIT 2", "SELECT count(*) FROM items WHERE vec LIKE json_vector('[1,2,3]') LIMIT 2", "SELECT id FROM items WHERE vec LIKE json_vector('[1]') LIMIT 2"} { + if _, e := c.Execute(context.Background(), q); e == nil { + t.Fatal(q) + } + } + p := &querySearchParams{values: map[string]interface{}{}} + p.AddRadius(1) + p.AddRangeFilter(2) + if p.Params()["radius"] != float64(1) { + t.Fatal(p) + } +} diff --git a/pkg/mysql.go b/pkg/mysql.go index 6e1b9b9..e82ce43 100644 --- a/pkg/mysql.go +++ b/pkg/mysql.go @@ -32,7 +32,11 @@ func (s *Server) serveMySQL(ctx context.Context, co net.Conn) { if err != nil { return } - defer conn.Close() + defer func() { + if conn.Conn != nil { + conn.Close() + } + }() co.SetDeadline(time.Time{}) for { if err := conn.HandleCommand(); err != nil { From dcebd167dfa27b21cf7fd55d578494d5d859a378 Mon Sep 17 00:00:00 2001 From: "shaoyue.chen" Date: Fri, 18 Sep 2026 20:19:13 +0800 Subject: [PATCH 3/5] Encode index parameters through SDK constructors and exercise all index methods --- pkg/commands.go | 51 ++++++++++++++++++++++++++++++++++++++++- pkg/integration_test.go | 17 +++++++++++++- 2 files changed, 66 insertions(+), 2 deletions(-) diff --git a/pkg/commands.go b/pkg/commands.go index a59d305..9267f69 100644 --- a/pkg/commands.go +++ b/pkg/commands.go @@ -4,6 +4,7 @@ import ( "encoding/json" "fmt" "regexp" + "strconv" "strings" "github.com/milvus-io/milvus-sdk-go/v2/client" @@ -128,7 +129,11 @@ func (c *ClientConn) handleMilvusCommand(sql string) (bool, error) { params["nlist"] = "128" } } - e := c.upstream.CreateIndex(c.ctx, m[2], m[3], sqlIndex{m[1], params["index_type"], params}, false, client.WithIndexName(m[1])) + idx, err := buildIndex(params) + if err != nil { + return true, err + } + e := c.upstream.CreateIndex(c.ctx, m[2], m[3], sqlIndex{m[1], params["index_type"], idx.Params()}, false, client.WithIndexName(m[1])) if e != nil { return true, e } @@ -197,3 +202,47 @@ func (c *ClientConn) handleMilvusCommand(sql string) (bool, error) { } return false, nil } + +func buildIndex(params map[string]string) (entity.Index, error) { + kind := params["index_type"] + metric := entity.MetricType(params["metric_type"]) + allowed := map[string]bool{"index_type": true, "metric_type": kind != "INVERTED"} + switch kind { + case "HNSW": + allowed["M"] = true + allowed["efConstruction"] = true + case "IVF_FLAT": + allowed["nlist"] = true + } + for key := range params { + if !allowed[key] { + return nil, fmt.Errorf("parameter %s is not valid for %s", key, kind) + } + } + switch kind { + case "FLAT": + return entity.NewIndexFlat(metric) + case "HNSW": + m, e := strconv.Atoi(params["M"]) + if e != nil { + return nil, e + } + ef, e := strconv.Atoi(params["efConstruction"]) + if e != nil { + return nil, e + } + return entity.NewIndexHNSW(metric, m, ef) + case "IVF_FLAT": + n, e := strconv.Atoi(params["nlist"]) + if e != nil { + return nil, e + } + return entity.NewIndexIvfFlat(metric, n) + case "AUTOINDEX": + return entity.NewIndexAUTOINDEX(metric) + case "INVERTED": + return entity.NewGenericIndex("", entity.IndexType(kind), map[string]string{"index_type": kind, "params": "{}"}), nil + default: + return nil, fmt.Errorf("unsupported index type %s", kind) + } +} diff --git a/pkg/integration_test.go b/pkg/integration_test.go index 8ba377f..b55a771 100644 --- a/pkg/integration_test.go +++ b/pkg/integration_test.go @@ -136,7 +136,7 @@ func TestMilvusIntegration(t *testing.T) { }() mustExec("USE " + dbname) mustExec("CREATE TABLE items (id bigint PRIMARY KEY, name varchar(100), enabled bool, score double, meta json, embedding vector(3))") - mustExec("CREATE INDEX embedding_idx ON items (embedding) USING FLAT WITH (metric_type='L2')") + mustExec("CREATE INDEX embedding_idx ON items (embedding) USING HNSW WITH (metric_type='L2', M=16, efConstruction=100)") mustExec("INSERT INTO items VALUES (1,'one',true,1.5,'{\"tag\":1}',json_vector('[1,0,0]')), (2,'two',false,2.5,'{\"tag\":2}',json_vector('[0,1,0]'))") mustExec("LOAD TABLE items") if v := mustQuery("SHOW INDEXES FROM items"); len(v) != 1 { @@ -181,6 +181,21 @@ func TestMilvusIntegration(t *testing.T) { mustExec("FLUSH TABLE items") mustExec("RELEASE TABLE items") mustExec("DROP INDEX embedding_idx ON items") + for _, method := range []string{"FLAT", "IVF_FLAT", "AUTOINDEX"} { + mustExec("CREATE INDEX embedding_idx ON items (embedding) USING " + method + " WITH (metric_type='L2')") + mustExec("LOAD TABLE items") + if v := mustQuery("SELECT id FROM items WHERE embedding LIKE json_vector('[1,0,0]') LIMIT 1"); len(v) != 1 { + t.Fatal(method, v) + } + mustExec("RELEASE TABLE items") + mustExec("DROP INDEX embedding_idx ON items") + } + mustExec("CREATE INDEX name_idx ON items (name) USING INVERTED") + if v := mustQuery("SHOW INDEXES FROM items"); len(v) != 1 { + t.Fatal(v) + } + mustExec("DROP INDEX name_idx ON items") + mustExec("CREATE PARTITION extra ON items") if v := mustQuery("SHOW PARTITIONS FROM items"); len(v) != 2 { t.Fatal(v) From d775393d79eee2cf764786a215b0f2f5bc2d9f09 Mon Sep 17 00:00:00 2001 From: "shaoyue.chen" Date: Fri, 18 Sep 2026 20:36:43 +0800 Subject: [PATCH 4/5] Harden protocol edge cases and enforce 90 percent coverage --- .github/workflows/test.yml | 5 +- Readme.md | 4 +- cmd/milvus-sql.go | 171 +++++++--------------- cmd/milvus-sql_test.go | 41 ++++++ pkg/ddl.go | 3 + pkg/encoding_test.go | 282 +++++++++++++++++++++++++++++++++++++ pkg/indexes.go | 26 ++++ pkg/integration_test.go | 22 ++- pkg/mysql.go | 4 +- pkg/postgres_state_test.go | 189 +++++++++++++++++++++++++ pkg/select.go | 10 +- pkg/select_plan.go | 2 +- 12 files changed, 624 insertions(+), 135 deletions(-) create mode 100644 cmd/milvus-sql_test.go create mode 100644 pkg/encoding_test.go create mode 100644 pkg/postgres_state_test.go diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index 916b5d3..dec0b43 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -13,7 +13,10 @@ jobs: with: go-version-file: go.mod cache: true - - run: go test -race ./... + - name: Unit tests and coverage gate + run: | + go test -race ./... -coverprofile=unit-coverage.out + go tool cover -func=unit-coverage.out | awk '/^total:/ { print; gsub("%", "", $3); if ($3 < 90) exit 1 }' - run: go vet ./... milvus: runs-on: ubuntu-latest diff --git a/Readme.md b/Readme.md index b843319..2edfd61 100644 --- a/Readme.md +++ b/Readme.md @@ -89,7 +89,7 @@ DROP DATABASE demo; Create does **not** implicitly index or load. Index methods: `FLAT`, `HNSW`, `IVF_FLAT`, `AUTOINDEX`, scalar `INVERTED`. Dense metrics: `L2` (default), `IP`, `COSINE`. Search uses the field's actual index metric, returns nearest neighbors -in Milvus order and optionally `_distance` (distance/similarity according to the +in Milvus order and optionally the reserved `_distance` (distance/similarity according to the metric). Only one positive vector predicate is allowed, optionally combined with scalar predicates using `AND`; vector predicates inside `OR` or `NOT` are rejected. Writes and searches currently target the default partition. @@ -118,7 +118,7 @@ go vet ./... MILVUS_TEST_ADDR=localhost:19530 go test -race ./pkg -run TestMilvusIntegration -v ``` -CI runs the lifecycle through both `database/sql` MySQL and pgx clients against a +CI requires at least 90% statement coverage and runs the lifecycle through both `database/sql` MySQL and pgx clients against a standalone Milvus container, including authentication failures and rejected SQL. Unit tests use isolated session mocks and real local protocol listeners. diff --git a/cmd/milvus-sql.go b/cmd/milvus-sql.go index d8f6e41..60cdea0 100644 --- a/cmd/milvus-sql.go +++ b/cmd/milvus-sql.go @@ -1,28 +1,13 @@ -// partially copied & changed from : https://github.com/flike/kingshard/blob/master/proxy/server/ - -// Copyright 2016 The kingshard Authors. All rights reserved. -// -// Licensed under the Apache License, Version 2.0 (the "License"): you may -// not use this file except in compliance with the License. You may obtain -// a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, software -// distributed under the License is distributed on an "AS IS" BASIS, WITHOUT -// WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. See the -// License for the specific language governing permissions and limitations -// under the License. - package main import ( + "context" "flag" "fmt" + "io" "os" "os/signal" - "path" - "runtime" + "path/filepath" "strings" "syscall" @@ -30,135 +15,73 @@ import ( "github.com/haorenfsa/milvus-sql-proxy/pkg" ) -var configFile *string = flag.String("config", "./config.yaml", "config file") -var logLevel *string = flag.String("log-level", "info", "log level [debug|info|warn|error], default info") -var version *bool = flag.Bool("v", false, "the version of kingshard") - -const ( - sqlLogName = "sql.log" - sysLogName = "sys.log" - MaxLogSize = 1024 * 1024 * 1024 -) - -var ( - BuildDate string - BuildVersion string -) - -const banner string = ` -████ ████ ██ ██ ████████ ███████ ██ -░██░██ ██░██░░ ░██ ██░░░░░░ ██░░░░░██ ░██ -░██░░██ ██ ░██ ██ ░██ ██ ██ ██ ██ ██████ ░██ ██ ░░██ ░██ -░██ ░░███ ░██░██ ░██░██ ░██░██ ░██ ██░░░░ ░█████████░██ ░██ ░██ -░██ ░░█ ░██░██ ░██░░██ ░██ ░██ ░██░░█████ ░░░░░░░░██░██ ██░██ ░██ -░██ ░ ░██░██ ░██ ░░████ ░██ ░██ ░░░░░██ ░██░░██ ░░ ██ ░██ -░██ ░██░██ ███ ░░██ ░░██████ ██████ ████████ ░░███████ ██░████████ -░░ ░░ ░░ ░░░ ░░ ░░░░░░ ░░░░░░ ░░░░░░░░ ░░░░░░░ ░░ ░░░░░░░░ -` +var BuildDate, BuildVersion string func main() { - fmt.Print(banner) - runtime.GOMAXPROCS(runtime.NumCPU()) - flag.Parse() - fmt.Printf("Git commit:%s\n", BuildVersion) - fmt.Printf("Build time:%s\n", BuildDate) - if *version { - return + ctx, cancel := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM, syscall.SIGQUIT) + defer cancel() + if err := run(ctx, os.Args[1:], os.Stderr); err != nil { + fmt.Fprintln(os.Stderr, err) + os.Exit(1) } - if len(*configFile) == 0 { - fmt.Println("must use a config file") - return +} +func run(ctx context.Context, args []string, out io.Writer) error { + flags := flag.NewFlagSet("milvus-sql-proxy", flag.ContinueOnError) + flags.SetOutput(out) + config := flags.String("config", "./config.yaml", "configuration file") + logLevel := flags.String("log-level", "", "override log level: debug, info, warn, error") + version := flags.Bool("v", false, "print build version") + if err := flags.Parse(args); err != nil { + return err } - - cfg, err := pkg.ParseConfigFile(*configFile) + if *version { + fmt.Fprintf(out, "milvus-sql-proxy %s (%s)\n", BuildVersion, BuildDate) + return nil + } + cfg, err := pkg.ParseConfigFile(*config) if err != nil { - fmt.Printf("parse config file error:%v\n", err.Error()) - return + return fmt.Errorf("configuration: %w", err) } - - //when the log file size greater than 1GB, kingshard will generate a new file - if len(cfg.LogPath) != 0 { - sysFilePath := path.Join(cfg.LogPath, sysLogName) - sysFile, err := golog.NewRotatingFileHandler(sysFilePath, MaxLogSize, 1) - if err != nil { - fmt.Printf("new log file error:%v\n", err.Error()) - return - } - golog.GlobalSysLogger = golog.New(sysFile, golog.Lfile|golog.Ltime|golog.Llevel) - - sqlFilePath := path.Join(cfg.LogPath, sqlLogName) - sqlFile, err := golog.NewRotatingFileHandler(sqlFilePath, MaxLogSize, 1) + level := cfg.LogLevel + if *logLevel != "" { + level = *logLevel + } + setLogLevel(level) + if cfg.LogPath != "" { + h, err := golog.NewRotatingFileHandler(filepath.Join(cfg.LogPath, "sys.log"), 1<<30, 1) if err != nil { - fmt.Printf("new log file error:%v\n", err.Error()) - return + return err } - golog.GlobalSqlLogger = golog.New(sqlFile, golog.Lfile|golog.Ltime|golog.Llevel) - } - - if *logLevel != "" { - setLogLevel(*logLevel) - } else { - setLogLevel(cfg.LogLevel) + old := golog.GlobalSysLogger + golog.GlobalSysLogger = golog.New(h, golog.Lfile|golog.Ltime|golog.Llevel) + defer func() { golog.GlobalSysLogger.Close(); golog.GlobalSysLogger = old }() + setLogLevel(level) } - golog.Info("main", "main", "log level is", 0, "level", *logLevel) - - var svr *pkg.Server - // var prometheusSvr *monitor.Prometheus - svr, err = pkg.NewServer(cfg) + s, err := pkg.NewServer(cfg) if err != nil { - golog.Error("main", "main", err.Error(), 0) - golog.GlobalSysLogger.Close() - golog.GlobalSqlLogger.Close() - return + return err } - // prometheusSvr, err = monitor.NewPrometheus(cfg.PrometheusAddr, svr) - // if err != nil { - // golog.Error("main", "main", err.Error(), 0) - // golog.GlobalSysLogger.Close() - // golog.GlobalSqlLogger.Close() - // svr.Close() - // return - // } - - sc := make(chan os.Signal, 1) - signal.Notify(sc, - syscall.SIGINT, - syscall.SIGTERM, - syscall.SIGQUIT, - syscall.SIGPIPE, - // syscall.SIGUSR1, - ) - + defer s.Close() + done := make(chan struct{}) + defer close(done) go func() { - for { - sig := <-sc - switch sig { - case syscall.SIGPIPE: - golog.Info("main", "main", "Ignore broken pipe signal", 0) - default: - golog.Info("main", "main", "Got signal", 0, "signal", sig) - golog.GlobalSysLogger.Close() - golog.GlobalSqlLogger.Close() - svr.Close() - } + select { + case <-ctx.Done(): + s.Close() + case <-done: } }() - // go prometheusSvr.Run() - golog.Info("main", "main", "starting server", 0) - svr.Run() + return s.Run() } - func setLogLevel(level string) { switch strings.ToLower(level) { case "debug": golog.GlobalSysLogger.SetLevel(golog.LevelDebug) - case "info": - golog.GlobalSysLogger.SetLevel(golog.LevelInfo) case "warn": golog.GlobalSysLogger.SetLevel(golog.LevelWarn) case "error": golog.GlobalSysLogger.SetLevel(golog.LevelError) default: - golog.GlobalSysLogger.SetLevel(golog.LevelError) + golog.GlobalSysLogger.SetLevel(golog.LevelInfo) } } diff --git a/cmd/milvus-sql_test.go b/cmd/milvus-sql_test.go new file mode 100644 index 0000000..f29811c --- /dev/null +++ b/cmd/milvus-sql_test.go @@ -0,0 +1,41 @@ +package main + +import ( + "bytes" + "context" + "os" + "path/filepath" + "testing" +) + +func TestRun(t *testing.T) { + var out bytes.Buffer + if e := run(context.Background(), []string{"-v"}, &out); e != nil || out.Len() == 0 { + t.Fatal(e) + } + if e := run(context.Background(), []string{"-bogus"}, &out); e == nil { + t.Fatal("bad flags accepted") + } + if e := run(context.Background(), []string{"-config", "missing"}, &out); e == nil { + t.Fatal("bad config accepted") + } + dir := t.TempDir() + config := filepath.Join(dir, "config.yaml") + if e := os.WriteFile(config, []byte("addr: 127.0.0.1:0\nlogPath: "+dir+"\n"), 0600); e != nil { + t.Fatal(e) + } + ctx, cancel := context.WithCancel(context.Background()) + cancel() + if e := run(ctx, []string{"-config", config, "-log-level", "debug"}, &out); e != nil { + t.Fatal(e) + } + if e := os.WriteFile(config, []byte("addr: invalid-address\n"), 0600); e != nil { + t.Fatal(e) + } + if e := run(ctx, []string{"-config", config}, &out); e == nil { + t.Fatal("invalid listener accepted") + } + for _, level := range []string{"debug", "warn", "error", "info", ""} { + setLogLevel(level) + } +} diff --git a/pkg/ddl.go b/pkg/ddl.go index 3fc5c0a..fc2689a 100644 --- a/pkg/ddl.go +++ b/pkg/ddl.go @@ -54,6 +54,9 @@ func DDLToMilvusSchema(stmt *sqlparser.DDL) (*MilvusSchema, error) { return nil, errors.New("duplicate field") } seen[f.Name] = true + if f.Name == "_distance" { + return nil, errors.New("_distance is reserved for vector search scores") + } if !fieldIdentifier.MatchString(f.Name) { return nil, errors.New("invalid field name") } diff --git a/pkg/encoding_test.go b/pkg/encoding_test.go new file mode 100644 index 0000000..abf863b --- /dev/null +++ b/pkg/encoding_test.go @@ -0,0 +1,282 @@ +package pkg + +import ( + "context" + "encoding/binary" + "errors" + "math" + "os" + "path/filepath" + "reflect" + "strings" + "testing" + "time" + + "github.com/flike/kingshard/mysql" + "github.com/milvus-io/milvus-proto/go-api/v2/commonpb" + "github.com/milvus-io/milvus-proto/go-api/v2/milvuspb" + "github.com/milvus-io/milvus-sdk-go/v2/client" + "github.com/milvus-io/milvus-sdk-go/v2/entity" + "google.golang.org/grpc" +) + +func TestScalarSerialization(t *testing.T) { + for _, tc := range []struct { + value interface{} + text string + }{ + {nil, "NULL"}, {true, "true"}, {false, "false"}, {int8(-8), "-8"}, {int16(-16), "-16"}, {int32(-32), "-32"}, {int64(-64), "-64"}, {int(-1), "-1"}, {uint8(8), "8"}, {uint16(16), "16"}, {uint32(32), "32"}, {uint64(64), "64"}, {uint(1), "1"}, {float32(1.5), "1.5"}, {float64(2.5), "2.5"}, {[]float32{1, 2}, "[1,2]"}, {[]byte("bytes"), "bytes"}, {"str", "str"}, + } { + b, e := formatValue(tc.value) + if e != nil || string(b) != tc.text { + t.Fatal(tc, b, e) + } + f := &mysql.Field{} + if e = formatField(f, tc.value); e != nil { + t.Fatal(tc, e) + } + } + if _, e := formatValue(struct{}{}); e == nil { + t.Fatal("accepted unsupported value") + } + if e := formatField(&mysql.Field{}, struct{}{}); e == nil { + t.Fatal("accepted unsupported field") + } + c := NewSession(context.Background(), nil) + if _, e := c.buildResultset(nil, []string{"a"}, [][]interface{}{{1, 2}}); e == nil { + t.Fatal("row size accepted") + } + if _, e := c.buildResultset([]*mysql.Field{{}}, []string{"a", "b"}, nil); e == nil { + t.Fatal("schema size accepted") + } + if _, e := c.buildResultset(nil, []string{"a"}, [][]interface{}{{struct{}{}}}); e == nil { + t.Fatal("value accepted") + } + for _, v := range []interface{}{nil, []byte("x"), true, false, int(1), uint64(3), float64(1.25)} { + if _, e := sqlLiteral(v); e != nil { + t.Fatal(e) + } + } + if _, e := sqlLiteral(struct{}{}); e == nil { + t.Fatal("parameter accepted") + } +} +func TestParameterInference(t *testing.T) { + c := NewSession(context.Background(), &operationMock{schema: testSchema()}) + for _, tc := range []struct { + q string + want []uint32 + }{ + {"SELECT id FROM items WHERE $1=id", []uint32{20}}, + {"SELECT id FROM items WHERE id IN ($1,$2) LIMIT $3 OFFSET $4", []uint32{20, 20, 20, 20}}, + {"DELETE FROM items WHERE enabled=$1", []uint32{16}}, + {"INSERT INTO items VALUES ($1,$2,$3,$4,$5,json_vector($6))", []uint32{20, 25, 16, 701, 25, 25}}, + {"INSERT INTO items(score,id) VALUES ($1,$2)", []uint32{701, 20}}, + {"SELECT $1", []uint32{25}}, + {"SHOW $1", []uint32{25}}, + } { + oids, e := c.parameterOIDs(context.Background(), tc.q, make([]uint32, len(tc.want))) + if e != nil || !reflect.DeepEqual(oids, tc.want) { + t.Fatal(tc.q, oids, e) + } + } + if v, e := c.parameterOIDs(context.Background(), "SELECT id FROM items WHERE id=$1", []uint32{23}); e != nil || v[0] != 23 { + t.Fatal(v, e) + } + c.upstream = &operationMock{fail: errors.New("schema unavailable")} + if _, e := c.parameterOIDs(context.Background(), "SELECT id FROM items WHERE id=$1", []uint32{0}); e == nil { + t.Fatal("schema failure ignored") + } +} +func TestParameterWireTypes(t *testing.T) { + for _, tc := range []struct { + oid uint32 + text string + want interface{} + }{{16, "true", true}, {20, "-5", int64(-5)}, {701, "1.5", float64(1.5)}, {25, "str", "str"}} { + v, e := pgParameter([]byte(tc.text), tc.oid, 0) + if e != nil || v != tc.want { + t.Fatal(v, e) + } + } + small := []byte{0xff, 0xfe} + if v, e := pgParameter(small, 21, 1); e != nil || v != int16(-2) { + t.Fatal(v, e) + } + real := make([]byte, 4) + binary.BigEndian.PutUint32(real, math.Float32bits(1.5)) + if v, e := pgParameter(real, 700, 1); e != nil || v != float32(1.5) { + t.Fatal(v, e) + } + if _, e := pgParameter([]byte{2}, 16, 1); e == nil { + t.Fatal("bad bool accepted") + } + if _, e := pgParameter([]byte("x"), 9999, 1); e == nil { + t.Fatal("unsupported OID accepted") + } + if v, e := pgValue(nil, 25, 1); e != nil || v != nil { + t.Fatal(v, e) + } + for _, oid := range []uint32{16, 20, 701} { + if _, e := pgValue("not numeric", oid, 1); e == nil { + t.Fatal(oid) + } + } + if formatAt([]int16{0, 1}, 3) != 0 { + t.Fatal("format overflow") + } +} + +type indexService struct { + milvuspb.MilvusServiceClient + response *milvuspb.DescribeIndexResponse + err error + request *milvuspb.DescribeIndexRequest +} + +func (s *indexService) DescribeIndex(_ context.Context, r *milvuspb.DescribeIndexRequest, _ ...grpc.CallOption) (*milvuspb.DescribeIndexResponse, error) { + s.request = r + return s.response, s.err +} +func TestIndexResponseFiltering(t *testing.T) { + service := &indexService{response: &milvuspb.DescribeIndexResponse{Status: &commonpb.Status{}, IndexDescriptions: []*milvuspb.IndexDescription{{IndexName: "scalar", FieldName: "name", Params: entity.MapKvPairs(map[string]string{"index_type": "INVERTED"})}, {IndexName: "first", FieldName: "v1", Params: entity.MapKvPairs(map[string]string{"index_type": "FLAT", "metric_type": "L2"})}, {IndexName: "second", FieldName: "v2", Params: entity.MapKvPairs(map[string]string{"index_type": "HNSW", "metric_type": "IP"})}}}} + c := NewSession(context.Background(), &client.GrpcClient{Service: service}) + indexes, e := c.vectorIndexes("items", "v2") + if e != nil || len(indexes) != 1 || indexes[0].Params()["metric_type"] != "IP" { + t.Fatal(indexes, e) + } + if service.request.FieldName != "v2" { + t.Fatal(service.request) + } + all, e := c.listIndexes("items") + if e != nil || len(all) != 3 { + t.Fatal(all, e) + } + if _, e := c.Execute(context.Background(), "SHOW INDEXES FROM items"); e != nil { + t.Fatal(e) + } + if _, e := c.Describe(context.Background(), "SHOW INDEXES FROM items"); e != nil { + t.Fatal(e) + } + service.response.Status = &commonpb.Status{ErrorCode: commonpb.ErrorCode_IndexNotExist} + all, e = c.listIndexes("items") + if e != nil || len(all) != 0 { + t.Fatal(all, e) + } + service.response.Status = &commonpb.Status{ErrorCode: commonpb.ErrorCode_UnexpectedError, Reason: "failure"} + if _, e = c.listIndexes("items"); e == nil { + t.Fatal("error lost") + } + if _, e = c.vectorIndexes("items", "v2"); e == nil { + t.Fatal("error lost") + } + service.err = errors.New("transport failed") + if _, e = c.listIndexes("items"); e == nil { + t.Fatal("error lost") + } + if _, e = c.vectorIndexes("items", "v2"); e == nil { + t.Fatal("error lost") + } +} +func TestIndexOptions(t *testing.T) { + for _, p := range []map[string]string{{"index_type": "HNSW", "M": "oops", "efConstruction": "100"}, {"index_type": "HNSW", "M": "16", "efConstruction": "oops"}, {"index_type": "IVF_FLAT", "nlist": "oops"}, {"index_type": "FLAT", "nlist": "100"}, {"index_type": "unsupported"}} { + if _, e := buildIndex(p); e == nil { + t.Fatal(p) + } + } + if _, e := buildIndex(map[string]string{"index_type": "AUTOINDEX", "metric_type": "L2"}); e != nil { + t.Fatal(e) + } +} +func TestConfigFiles(t *testing.T) { + path := filepath.Join(t.TempDir(), "config.yaml") + if e := os.WriteFile(path, []byte("mode: both\n"), 0600); e != nil { + t.Fatal(e) + } + if _, e := ParseConfigFile(path); e != nil { + t.Fatal(e) + } + if _, e := ParseConfigFile(path + ".missing"); e == nil { + t.Fatal("missing config accepted") + } + if _, e := ParseConfigData([]byte("[bad")); e == nil { + t.Fatal("invalid yaml accepted") + } +} + +func TestIndexAndSchemaFailures(t *testing.T) { + m := &operationMock{schema: testSchema()} + c := NewSession(context.Background(), m) + for _, q := range []string{"", "SELECT (", "SELECT COUNT(*) FROM items LIMIT 1", "SELECT id FROM items LIMIT 16385", "SELECT unknown FROM items", "CREATE TABLE bad (id bigint primary key, _distance bigint, v vector(3))", "CREATE TABLE bad (id bigint primary key, v vector)", "CREATE TABLE bad (id bigint primary key, name varchar, v vector(3))", "CREATE TABLE bad (id bigint primary key, x bigint auto_increment, v vector(3))", "CREATE TABLE bad (id bigint primary key, x bigint unique, v vector(3))", "DELETE FROM items i WHERE id=1", "DELETE FROM items,other WHERE id=1", "DELETE FROM items WHERE id IS NULL"} { + if _, e := c.Execute(context.Background(), q); e == nil { + t.Fatal(q) + } + } + if _, e := c.Execute(context.Background(), "SELECT COUNT(*) FROM items"); e != nil { + t.Fatal(e) + } + m.schema = nil + if _, e := c.GetCollectinSchema("items"); e == nil { + t.Fatal("missing schema accepted") + } + m.schema = testSchema() + m.fail = errors.New("denied") + if e := c.handleUseDB("other", nil); e == nil { + t.Fatal("database error ignored") + } + if _, e := c.Execute(context.Background(), "CREATE TABLE t (id bigint primary key,v vector(3))"); e == nil { + t.Fatal("create error ignored") + } + if _, e := emptyColumn(&entity.Field{DataType: entity.FieldTypeFloatVector}); e == nil { + t.Fatal("dimension missing") + } + if _, e := emptyColumn(&entity.Field{DataType: entity.FieldTypeBinaryVector}); e == nil { + t.Fatal("unsupported type accepted") + } +} +func TestMySQLResultMetadata(t *testing.T) { + c := NewSession(context.Background(), nil) + r, e := c.buildResultset(nil, []string{"flag", "vec"}, [][]interface{}{{false, []float32{1, 2}}}) + if e != nil { + t.Fatal(e) + } + if _, e = mysqlResult(&mysql.Result{Resultset: r}, false); e != nil { + t.Fatal(e) + } + empty := newEmptyResultset([]string{"id"}) + empty.Fields[0] = resultField("id", entity.FieldTypeInt64) + mr, e := mysqlResult(&mysql.Result{Resultset: empty}, true) + if e != nil || mr.Fields[0].Type != mysql.MYSQL_TYPE_LONGLONG { + t.Fatal(mr, e) + } + h := &mysqlHandler{ctx: context.Background(), timeout: time.Second, session: NewSession(context.Background(), &operationMock{schema: testSchema()})} + if _, e := h.HandleFieldList("items", ""); e == nil { + t.Fatal("field list accepted") + } + if e := h.HandleOtherCommand(0, nil); e == nil { + t.Fatal("command accepted") + } + if _, _, _, e := h.HandleStmtPrepare("SELECT '"); e == nil { + t.Fatal("bad prepare") + } + if _, e := h.HandleStmtExecute(nil, "SELECT id FROM items WHERE id=?", nil); e == nil { + t.Fatal("missing parameter accepted") + } +} +func TestPostgresCommandTags(t *testing.T) { + for _, tc := range []struct{ q, want string }{{"INSERT INTO x VALUES (1)", "INSERT 0 2"}, {"UPSERT INTO x VALUES (1)", "INSERT 0 2"}, {"DELETE FROM x WHERE id=1", "DELETE 2"}, {"CREATE TABLE x", "CREATE TABLE"}, {"DROP DATABASE x", "DROP DATABASE"}, {"LOAD TABLE x", "LOAD"}, {"", ""}} { + if v := pgTag(tc.q, &mysql.Result{AffectedRows: 2}); v != tc.want { + t.Fatal(v, tc.want) + } + } +} +func TestNormalizeTypesQuotedValues(t *testing.T) { + q := "CREATE TABLE `a` (`id` BIGINT PRIMARY KEY, name VARCHAR(9) COMMENT 'bool \\' x', v VECTOR(3), flag BOOLEAN)" + got := normalizeTypes(q) + if !strings.Contains(got, "flag tinyint(1)") || !strings.Contains(got, "COMMENT 'bool \\' x'") { + t.Fatal(got) + } + if normalizeTypes("SELECT 'boolean'") != "SELECT 'boolean'" { + t.Fatal("rewrote non-DDL") + } +} diff --git a/pkg/indexes.go b/pkg/indexes.go index f1cc217..eff7c92 100644 --- a/pkg/indexes.go +++ b/pkg/indexes.go @@ -33,3 +33,29 @@ func (c *ClientConn) listIndexes(table string) ([]entity.Index, error) { } return indexes, nil } + +// Recent Milvus versions may return multiple fields' indexes for the legacy +// DescribeIndex request. Match the response field explicitly before choosing a +// metric; otherwise scalar/other vector indexes can change search semantics. +func (c *ClientConn) vectorIndexes(table, field string) ([]entity.Index, error) { + if upstream, ok := c.upstream.(*client.GrpcClient); ok { + resp, err := upstream.Service.DescribeIndex(c.ctx, &milvuspb.DescribeIndexRequest{CollectionName: table, FieldName: field}) + if err != nil { + return nil, err + } + status := resp.GetStatus() + if status == nil || status.GetErrorCode() != commonpb.ErrorCode_Success || status.GetCode() != 0 { + return nil, fmt.Errorf("describe vector index: %s", status.GetReason()) + } + indexes := []entity.Index{} + for _, d := range resp.GetIndexDescriptions() { + if d.FieldName != field { + continue + } + params := entity.KvPairsMap(d.Params) + indexes = append(indexes, entity.NewGenericIndex(d.IndexName, entity.IndexType(params["index_type"]), params)) + } + return indexes, nil + } + return c.upstream.DescribeIndex(c.ctx, table, field) +} diff --git a/pkg/integration_test.go b/pkg/integration_test.go index b55a771..0e2053b 100644 --- a/pkg/integration_test.go +++ b/pkg/integration_test.go @@ -128,6 +128,7 @@ func TestMilvusIntegration(t *testing.T) { defer func() { exec("USE " + dbname) exec("RELEASE TABLE items") + exec("DROP TABLE multi") exec("DROP TABLE items") exec("USE default") if e := exec("DROP DATABASE " + dbname); e != nil { @@ -137,9 +138,10 @@ func TestMilvusIntegration(t *testing.T) { mustExec("USE " + dbname) mustExec("CREATE TABLE items (id bigint PRIMARY KEY, name varchar(100), enabled bool, score double, meta json, embedding vector(3))") mustExec("CREATE INDEX embedding_idx ON items (embedding) USING HNSW WITH (metric_type='L2', M=16, efConstruction=100)") + mustExec("CREATE INDEX name_idx ON items (name) USING INVERTED") mustExec("INSERT INTO items VALUES (1,'one',true,1.5,'{\"tag\":1}',json_vector('[1,0,0]')), (2,'two',false,2.5,'{\"tag\":2}',json_vector('[0,1,0]'))") mustExec("LOAD TABLE items") - if v := mustQuery("SHOW INDEXES FROM items"); len(v) != 1 { + if v := mustQuery("SHOW INDEXES FROM items"); len(v) != 2 { t.Fatalf("indexes: %v", v) } if v := mustQuery("DESCRIBE items"); len(v) != 6 { @@ -151,7 +153,7 @@ func TestMilvusIntegration(t *testing.T) { if v := mustQuery("SELECT id FROM items WHERE id=999"); len(v) != 0 { t.Fatal(v) } - if v := mustQuery("SELECT count(*) FROM items"); len(v) != 1 { + if v := mustQuery("SELECT COUNT(*) FROM items"); len(v) != 1 || fmt.Sprint(v[0][0]) != "2" { t.Fatal(v) } if v := mustQuery("SELECT id,_distance FROM items WHERE embedding LIKE json_vector('[1,0,0]') LIMIT 1"); len(v) != 1 || fmt.Sprint(v[0][0]) != "1" { @@ -164,8 +166,11 @@ func TestMilvusIntegration(t *testing.T) { if v := mustQuery("SELECT id,name FROM items WHERE id="+placeholder, int64(2)); len(v) != 1 { t.Fatalf("parameter query: %v", v) } + if v := mustQuery("SELECT id,name FROM items WHERE id="+placeholder, int64(999)); len(v) != 0 { + t.Fatalf("empty prepared query: %v", v) + } mustExec("UPSERT INTO items VALUES (2,'updated',true,3.5,'{}',json_vector('[0,0,1]'))") - if v := mustQuery("SELECT name FROM items WHERE id=2"); len(v) != 1 { + if v := mustQuery("SELECT name FROM items WHERE id=2"); len(v) != 1 || fmt.Sprint(v[0][0]) != "updated" { t.Fatal(v) } mustExec("DELETE FROM items WHERE id=2") @@ -180,6 +185,7 @@ func TestMilvusIntegration(t *testing.T) { mustQuery("SELECT id FROM items WHERE id=1") mustExec("FLUSH TABLE items") mustExec("RELEASE TABLE items") + mustExec("DROP INDEX name_idx ON items") mustExec("DROP INDEX embedding_idx ON items") for _, method := range []string{"FLAT", "IVF_FLAT", "AUTOINDEX"} { mustExec("CREATE INDEX embedding_idx ON items (embedding) USING " + method + " WITH (metric_type='L2')") @@ -202,6 +208,16 @@ func TestMilvusIntegration(t *testing.T) { } mustExec("DROP PARTITION extra ON items") mustExec("DROP TABLE items") + mustExec("CREATE TABLE multi (id BIGINT PRIMARY KEY, v1 VECTOR(3), v2 VECTOR(3))") + mustExec("CREATE INDEX i1 ON multi (v1) USING FLAT WITH (metric_type='L2')") + mustExec("CREATE INDEX i2 ON multi (v2) USING FLAT WITH (metric_type='IP')") + mustExec("INSERT INTO multi VALUES (1,json_vector('[1,0,0]'),json_vector('[0,1,0]')),(2,json_vector('[0,1,0]'),json_vector('[1,0,0]'))") + mustExec("LOAD TABLE multi") + if v := mustQuery("SELECT id FROM multi WHERE v2 LIKE json_vector('[1,0,0]') LIMIT 1"); len(v) != 1 || fmt.Sprint(v[0][0]) != "2" { + t.Fatalf("multi-vector index selection: %v", v) + } + mustExec("DROP TABLE multi") + }) } // Incorrect credentials must fail both protocol handshakes. diff --git a/pkg/mysql.go b/pkg/mysql.go index e82ce43..08e5839 100644 --- a/pkg/mysql.go +++ b/pkg/mysql.go @@ -94,8 +94,8 @@ func mysqlResult(r *legacy.Result, binary bool) (*mysql.Result, error) { // Preserve field metadata for empty result sets as well as nonempty rows. for i, f := range r.Fields { if len(values) == 0 { - rs.Fields[i].Type = f.Type - rs.Fields[i].Charset = f.Charset + rs.Fields[i] = &mysql.Field{Name: append([]byte(nil), f.Name...), Type: f.Type, Charset: f.Charset} + rs.FieldNames[string(f.Name)] = i } } out.Resultset = rs diff --git a/pkg/postgres_state_test.go b/pkg/postgres_state_test.go new file mode 100644 index 0000000..64fa993 --- /dev/null +++ b/pkg/postgres_state_test.go @@ -0,0 +1,189 @@ +package pkg + +import ( + "context" + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "crypto/tls" + "crypto/x509" + "crypto/x509/pkix" + "database/sql" + "encoding/pem" + "fmt" + "math/big" + "net" + "os" + "path/filepath" + "testing" + "time" + + driver "github.com/go-sql-driver/mysql" + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgproto3" +) + +func rawPG(t *testing.T, s *Server) *pgproto3.Frontend { + t.Helper() + co, e := net.Dial("tcp", s.listeners[1].Addr().String()) + if e != nil { + t.Fatal(e) + } + t.Cleanup(func() { co.Close() }) + co.SetDeadline(time.Now().Add(10 * time.Second)) + f := pgproto3.NewFrontend(co, co) + f.Send(&pgproto3.StartupMessage{ProtocolVersion: 196608, Parameters: map[string]string{"user": "root", "database": "default"}}) + if e = f.Flush(); e != nil { + t.Fatal(e) + } + if _, e = f.Receive(); e != nil { + t.Fatal(e) + } + f.Send(&pgproto3.PasswordMessage{Password: "secret"}) + f.Flush() + for { + msg, e := f.Receive() + if e != nil { + t.Fatal(e) + } + if _, ok := msg.(*pgproto3.ReadyForQuery); ok { + break + } + } + return f +} +func TestPostgresProtocolErrors(t *testing.T) { + s := mockServer(t) + for _, tc := range []struct { + name string + messages []pgproto3.FrontendMessage + bad bool + }{ + {"missing statement", []pgproto3.FrontendMessage{&pgproto3.Bind{PreparedStatement: "absent"}}, true}, + {"missing describe statement", []pgproto3.FrontendMessage{&pgproto3.Describe{ObjectType: 'S', Name: "absent"}}, true}, + {"missing describe portal", []pgproto3.FrontendMessage{&pgproto3.Describe{ObjectType: 'P', Name: "absent"}}, true}, + {"bad describe kind", []pgproto3.FrontendMessage{&pgproto3.Describe{ObjectType: 'Z'}}, true}, + {"missing execute portal", []pgproto3.FrontendMessage{&pgproto3.Execute{Portal: "absent"}}, true}, + {"invalid close kind", []pgproto3.FrontendMessage{&pgproto3.Close{ObjectType: 'Z'}}, true}, + {"malformed parse", []pgproto3.FrontendMessage{&pgproto3.Parse{Name: "a", Query: "select 'oops"}}, true}, + {"too many type hints", []pgproto3.FrontendMessage{&pgproto3.Parse{Name: "a", Query: "select 1", ParameterOIDs: []uint32{20}}}, true}, + {"duplicate statement", []pgproto3.FrontendMessage{&pgproto3.Parse{Name: "a", Query: "select 1"}, &pgproto3.Parse{Name: "a", Query: "select 1"}}, true}, + {"parameter count mismatch", []pgproto3.FrontendMessage{&pgproto3.Parse{Name: "a", Query: "select id from items where id=$1"}, &pgproto3.Bind{PreparedStatement: "a"}}, true}, + {"bad parameter format", []pgproto3.FrontendMessage{&pgproto3.Parse{Name: "a", Query: "select id from items where id=$1"}, &pgproto3.Bind{PreparedStatement: "a", Parameters: [][]byte{[]byte("1")}, ParameterFormatCodes: []int16{2}}}, true}, + {"bad result format", []pgproto3.FrontendMessage{&pgproto3.Parse{Name: "a", Query: "select id from items"}, &pgproto3.Bind{PreparedStatement: "a", ResultFormatCodes: []int16{2}}}, true}, + {"duplicate portal", []pgproto3.FrontendMessage{&pgproto3.Parse{Name: "a", Query: "select id from items"}, &pgproto3.Bind{PreparedStatement: "a", DestinationPortal: "p"}, &pgproto3.Bind{PreparedStatement: "a", DestinationPortal: "p"}}, true}, + {"mutation describe no data", []pgproto3.FrontendMessage{&pgproto3.Parse{Name: "a", Query: "create database demo"}, &pgproto3.Describe{ObjectType: 'S', Name: "a"}}, false}, + {"bind describe execute", []pgproto3.FrontendMessage{&pgproto3.Parse{Name: "a", Query: "select id from items"}, &pgproto3.Bind{PreparedStatement: "a", DestinationPortal: "p", ResultFormatCodes: []int16{1}}, &pgproto3.Describe{ObjectType: 'P', Name: "p"}, &pgproto3.Execute{Portal: "p", MaxRows: 1}, &pgproto3.Execute{Portal: "p"}, &pgproto3.Close{ObjectType: 'P', Name: "p"}, &pgproto3.Close{ObjectType: 'S', Name: "a"}, &pgproto3.Flush{}}, false}, + {"execution failure", []pgproto3.FrontendMessage{&pgproto3.Parse{Name: "a", Query: "delete from items"}, &pgproto3.Bind{PreparedStatement: "a"}, &pgproto3.Execute{}}, true}, + {"unsupported frontend message", []pgproto3.FrontendMessage{&pgproto3.CopyData{Data: []byte("bad")}}, true}, + } { + t.Run(tc.name, func(t *testing.T) { + f := rawPG(t, s) + for _, m := range tc.messages { + f.Send(m) + } + f.Send(&pgproto3.Sync{}) + if e := f.Flush(); e != nil { + t.Fatal(e) + } + bad := false + for { + m, e := f.Receive() + if e != nil { + t.Fatal(e) + } + if _, ok := m.(*pgproto3.ErrorResponse); ok { + bad = true + } + if _, ok := m.(*pgproto3.ReadyForQuery); ok { + break + } + } + if bad != tc.bad { + t.Fatalf("error=%v wanted %v", bad, tc.bad) + } + }) + } + for _, q := range []string{"", "select 'unterminated", "select nope from items"} { + f := rawPG(t, s) + f.Send(&pgproto3.Query{String: q}) + f.Flush() + for { + m, e := f.Receive() + if e != nil { + t.Fatal(e) + } + if _, ok := m.(*pgproto3.ReadyForQuery); ok { + break + } + } + } +} +func TestFrontendTLS(t *testing.T) { + key, e := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + if e != nil { + t.Fatal(e) + } + cert := &x509.Certificate{SerialNumber: big.NewInt(1), Subject: pkix.Name{CommonName: "sqlproxy test"}, NotBefore: time.Now().Add(-time.Hour), NotAfter: time.Now().Add(time.Hour), IPAddresses: []net.IP{net.ParseIP("127.0.0.1")}, KeyUsage: x509.KeyUsageDigitalSignature, ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}, BasicConstraintsValid: true} + der, e := x509.CreateCertificate(rand.Reader, cert, cert, &key.PublicKey, key) + if e != nil { + t.Fatal(e) + } + priv, e := x509.MarshalPKCS8PrivateKey(key) + if e != nil { + t.Fatal(e) + } + dir := t.TempDir() + certPath := filepath.Join(dir, "cert.pem") + keyPath := filepath.Join(dir, "key.pem") + certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der}) + if e = os.WriteFile(certPath, certPEM, 0600); e != nil { + t.Fatal(e) + } + if e = os.WriteFile(keyPath, pem.EncodeToMemory(&pem.Block{Type: "PRIVATE KEY", Bytes: priv}), 0600); e != nil { + t.Fatal(e) + } + s, e := NewServer(&Config{Mode: "both", Addr: "127.0.0.1:0", PostgresAddr: "127.0.0.1:0", User: "root", Password: "secret", TLSCert: certPath, TLSKey: keyPath}) + if e != nil { + t.Fatal(e) + } + s.newSession = func(ctx context.Context) (*ClientConn, error) { + return NewSession(ctx, &operationMock{schema: testSchema()}), nil + } + done := make(chan error, 1) + go func() { done <- s.Run() }() + defer func() { s.Close(); <-done }() + roots := x509.NewCertPool() + roots.AppendCertsFromPEM(certPEM) + if e = driver.RegisterTLSConfig("sqlproxy-test", &tls.Config{RootCAs: roots, MinVersion: tls.VersionTLS12}); e != nil { + t.Fatal(e) + } + defer driver.DeregisterTLSConfig("sqlproxy-test") + db, e := sql.Open("mysql", fmt.Sprintf("root:secret@tcp(%s)/default?tls=sqlproxy-test", s.listeners[0].Addr())) + if e != nil { + t.Fatal(e) + } + defer db.Close() + if e = db.Ping(); e != nil { + t.Fatal(e) + } + pg, e := pgx.Connect(context.Background(), fmt.Sprintf("postgres://root:secret@%s/default?sslmode=verify-full&sslrootcert=%s", s.listeners[1].Addr(), certPath)) + if e != nil { + t.Fatal(e) + } + defer pg.Close(context.Background()) + var n int + if e = pg.QueryRow(context.Background(), "SELECT 1").Scan(&n); e != nil || n != 1 { + t.Fatal(n, e) + } + // SSL negotiation falls back only when the client explicitly permits it. + plain := mockServer(t) + pg2, e := pgx.Connect(context.Background(), fmt.Sprintf("postgres://root:secret@%s/default?sslmode=prefer", plain.listeners[1].Addr())) + if e != nil { + t.Fatal(e) + } + pg2.Close(context.Background()) + if _, e := NewServer(&Config{TLSCert: "missing", TLSKey: "missing"}); e == nil { + t.Fatal("invalid TLS files accepted") + } +} diff --git a/pkg/select.go b/pkg/select.go index da2a226..3e9f184 100644 --- a/pkg/select.go +++ b/pkg/select.go @@ -38,6 +38,9 @@ func (c *ClientConn) handleSelect(stmt *sqlparser.Select, _ []interface{}) error pk = f.Name } } + if plan.vectorField != "" && schemaFields["_distance"] != nil { + return fmt.Errorf("_distance is reserved for vector search scores") + } names := []string{} types := []entity.FieldType{} if len(stmt.SelectExprs) == 1 && sqlparser.String(stmt.SelectExprs[0]) == "*" { @@ -48,6 +51,9 @@ func (c *ClientConn) handleSelect(stmt *sqlparser.Select, _ []interface{}) error } else { for _, expr := range stmt.SelectExprs { name := sqlparser.String(expr) + if strings.EqualFold(name, "count(*)") { + name = "count(*)" + } if name != "count(*)" { name = expr.(*sqlparser.AliasedExpr).Expr.(*sqlparser.ColName).Name.String() } @@ -89,7 +95,7 @@ func (c *ClientConn) handleSelect(stmt *sqlparser.Select, _ []interface{}) error outputs := []string{} seen := map[string]bool{} for _, n := range names { - if n != "_distance" && !seen[n] { + if (n != "_distance" || plan.vectorField == "") && !seen[n] { outputs = append(outputs, n) seen[n] = true } @@ -104,7 +110,7 @@ func (c *ClientConn) handleSelect(stmt *sqlparser.Select, _ []interface{}) error if e != nil { return e } - indexes, e := c.upstream.DescribeIndex(c.ctx, plan.table, plan.vectorField) + indexes, e := c.vectorIndexes(plan.table, plan.vectorField) if e != nil { return e } diff --git a/pkg/select_plan.go b/pkg/select_plan.go index c481f78..d3d5620 100644 --- a/pkg/select_plan.go +++ b/pkg/select_plan.go @@ -46,7 +46,7 @@ func planSelect(s *sqlparser.Select) (selectPlan, error) { if !ok || !a.As.IsEmpty() { return p, fmt.Errorf("projection aliases are not supported") } - if sqlparser.String(a.Expr) == "count(*)" && len(s.SelectExprs) == 1 { + if strings.EqualFold(sqlparser.String(a.Expr), "count(*)") && len(s.SelectExprs) == 1 { continue } c, ok := a.Expr.(*sqlparser.ColName) From d718f50bfd37a82ab6aa89358511e32ebcdcefb2 Mon Sep 17 00:00:00 2001 From: "shaoyue.chen" Date: Fri, 18 Sep 2026 20:38:40 +0800 Subject: [PATCH 5/5] Clarify partition scope for queries and mutations --- Readme.md | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/Readme.md b/Readme.md index 2edfd61..ff8e88a 100644 --- a/Readme.md +++ b/Readme.md @@ -92,7 +92,8 @@ Create does **not** implicitly index or load. Index methods: `FLAT`, `HNSW`, in Milvus order and optionally the reserved `_distance` (distance/similarity according to the metric). Only one positive vector predicate is allowed, optionally combined with scalar predicates using `AND`; vector predicates inside `OR` or `NOT` are rejected. -Writes and searches currently target the default partition. +Inserts/upserts target the default partition. Query/search/delete have no partition +selector; they follow Milvus behavior across the collection (searching loaded partitions). Scalar predicates: `=`, `!=`, `<>`, `<`, `<=`, `>`, `>=`, `IN`, `NOT IN`, `LIKE`, `AND`, `OR`, `NOT`, parentheses. `LIMIT count OFFSET offset` and `LIMIT offset,count`