diff --git a/.github/workflows/ci-build-tests.yml b/.github/workflows/ci-build-tests.yml index da009a021..d31311fa6 100644 --- a/.github/workflows/ci-build-tests.yml +++ b/.github/workflows/ci-build-tests.yml @@ -127,12 +127,14 @@ jobs: runs-on: ubuntu-latest env: TRICKSTER_MYSQL_CLI_TEST: "1" + TRICKSTER_GREPTIMEDB_MYSQL_ACCEPTANCE: "1" + TRICKSTER_GREPTIMEDB_PROMQL_ACCEPTANCE: "1" TRICKSTER_DNS_TEST: "1" # CI containers cannot raise the UDP receive buffer; the warning is noise QUIC_GO_DISABLE_RECEIVE_BUFFER_WARNING: "1" # a tenth of the synthetic trips rows: the integration tests only need # non-empty, correctly shaped data, and the database loads are the - # slowest part of bringing the environment up + # slowest part of bringing the environment up (including GreptimeDB) SEED_PROFILE: small steps: - uses: actions/checkout@v6 @@ -150,9 +152,20 @@ jobs: docker compose logs --tail=30 devorigin prometheus_seed_generate prometheus_seed docker compose logs --tail=60 graphite graphite_seed graphite_generator - name: Run integration tests (with coverage) + env: + GREPTIMEDB_REPORT_DIR: ${{ runner.temp }}/greptimedb-reports run: make integration-cover - name: Run integration tests (with race detection enabled) + env: + GREPTIMEDB_REPORT_DIR: ${{ runner.temp }}/greptimedb-reports run: make -C integration data-race-test + - name: Upload GreptimeDB acceptance reports + if: always() + uses: actions/upload-artifact@v4 + with: + name: greptimedb-acceptance-reports + path: ${{ runner.temp }}/greptimedb-reports + if-no-files-found: ignore - name: Send integration coverage if: github.repository_owner == 'trickstercache' uses: shogo82148/actions-goveralls@v1 diff --git a/.gitignore b/.gitignore index 4cf20a5ec..c16beee1d 100644 --- a/.gitignore +++ b/.gitignore @@ -78,4 +78,5 @@ AGENTS.md # (seeded from docker-compose-data/coredns/*.seed by `make integration-start`) docs/developer/environment/docker-compose-data/coredns-zones/ -integration/conformance/reports \ No newline at end of file +integration/conformance/reports +integration/greptimedb/reports diff --git a/.golangci.yml b/.golangci.yml index 8428668ef..3cba203ac 100644 --- a/.golangci.yml +++ b/.golangci.yml @@ -25,6 +25,7 @@ linters: - intrange - exptostd - errorlint + - forbidigo - gocheckcompilerdirectives - mirror - importas @@ -35,6 +36,13 @@ linters: settings: dupl: threshold: 150 # default is 100; have brought current codebase down to 75, but leaving default at 150 for now + forbidigo: + # Calls are centralized except for the standalone-core files excluded below. + analyze-types: true + forbid: + - pattern: ^rand\.[^.]+$ + pkg: ^math/rand(/v2)?$ + msg: use pkg/util/weak/compat in the application or pkg/util/weak/weaktest in tests lll: line-length: 180 staticcheck: @@ -159,6 +167,14 @@ linters: - common-false-positives - legacy - std-error-handling + rules: + - path: pkg/util/weak/ + linters: + - forbidigo + # Keep the standalone load-balancing core's existing, audited randomness. + - path: ^pkg/lb/(hrw/hrw|p2c/p2c|rr/rr|lbtest/lbtest)\.go$ + linters: + - forbidigo paths: - third_party$ - builtin$ diff --git a/Makefile b/Makefile index 80fe4ed0e..526f9b801 100644 --- a/Makefile +++ b/Makefile @@ -166,6 +166,12 @@ style: check-imports: @go run hack/check-imports/main.go +# fails the build if weak randomness crosses the application/test boundary; +# pkg/util/weak/weaktest also panics at runtime once the application registers +.PHONY: check-weak-random +check-weak-random: + @go run hack/check-weak-random/main.go + .PHONY: gofix-apply gofix-apply: @go fix ./... @@ -178,12 +184,12 @@ LINT_FLAGS ?= .PHONY: golangci-lint golangci-lint: @go tool golangci-lint run $(LINT_FLAGS) -c .golangci.yml - @for m in hack/seedgen hack/druidseed hack/devorigin; do \ + @for m in hack/seedgen hack/druidseed hack/greptimeseed hack/devorigin; do \ (cd $$m && go tool -modfile ../../go.mod golangci-lint run $(LINT_FLAGS) -c ../../.golangci.yml ./...) || exit 1; \ done .PHONY: lint -lint: check-imports spelling vulncheck gofix-diff golangci-lint +lint: check-imports check-weak-random spelling vulncheck gofix-diff golangci-lint .PHONY: lint-all lint-all: @@ -229,14 +235,14 @@ lint-fix: GO_TEST_FLAGS ?= -coverprofile=.coverprofile .PHONY: test -test: check-license-headers check-codegen gotest check-fmtprints check-todos check-devorigin-offline +test: check-license-headers check-codegen check-weak-random gotest check-fmtprints check-todos check-devorigin-offline GO_TEST_PATH ?= $(shell $(GO) list ./... | grep -v v2/integration | tr '\n' ' ') .PHONY: gotest gotest: $(GO) test -timeout=5m -v ${GO_TEST_FLAGS} $(GO_TEST_PATH) @./hack/filter-coverprofile.sh .coverprofile - @for m in hack/seedgen hack/druidseed hack/devorigin; do (cd $$m && $(GO) test -timeout=5m ./...) || exit 1; done + @for m in hack/seedgen hack/druidseed hack/greptimeseed hack/devorigin; do (cd $$m && $(GO) test -timeout=5m ./...) || exit 1; done @echo @./hack/coverprofile-summary.sh @echo "All tests passed successfully." @@ -556,6 +562,11 @@ seed-generate: @cd hack/seedgen && $(GO) run . -out ../../docs/developer/environment/docker-compose-data/seed-data \ $(if $(SEED_PROFILE),-profile $(SEED_PROFILE),) $(if $(SEED_FORCE),-force,) +# Read-only direct GreptimeDB acceptance; does not start or reseed services. +.PHONY: developer-greptimedb-check +developer-greptimedb-check: + @GO="$(GO)" sh hack/greptimedb-check.sh + RUN_FLAGS ?= .PHONY: serve-dev serve-dev: diff --git a/cmd/trickster/main.go b/cmd/trickster/main.go index 1aa7364c6..ebd426b6e 100644 --- a/cmd/trickster/main.go +++ b/cmd/trickster/main.go @@ -26,6 +26,7 @@ import ( "github.com/trickstercache/trickster/v2/pkg/observability/keys" "github.com/trickstercache/trickster/v2/pkg/observability/logging" "github.com/trickstercache/trickster/v2/pkg/observability/logging/logger" + "github.com/trickstercache/trickster/v2/pkg/util/weak" ) // application variables set at build time via go build's -ldflags @@ -36,6 +37,8 @@ var ( ) func main() { + // from here on, test-only randomness panics unless a test launched this process + weak.RegisterUnlessTestMode() appinfo.Set(appinfo.AppName, applicationVersion, applicationBuildTime, applicationGitCommitID) err := daemon.Start(context.Background(), os.Args[1:]...) diff --git a/deploy/kube/configmap.yaml b/deploy/kube/configmap.yaml index d6254b067..7086a0ca0 100644 --- a/deploy/kube/configmap.yaml +++ b/deploy/kube/configmap.yaml @@ -501,6 +501,19 @@ data: # # timeseries_retention_factor, backfill_tolerance, and # # max_object_size_bytes backend settings apply. + # # GreptimeDB serves HTTP and either or both native protocols. See + # # examples/conf/greptimedb.yaml for the corresponding listeners/authenticator. + # greptimedb1: + # provider: greptimedb + # origin_url: http://greptime.example:4000 + # listener_names: [default, greptime-pg, greptime-mysql] + # authenticator_name: greptime-readers + # cache_name: default + # postgres: + # upstream_url: postgres://readonly:REPLACE_ME@greptime.example:4003/public + # mysql: + # upstream_url: mysql://readonly:REPLACE_ME@greptime.example:4002/public + # # # example postgres backend, exposed by a listener with protocol postgres # # (see the listeners section). Provider timescaledb is an alias of postgres. # postgres1: diff --git a/docs/developer/adding-pgwire-engine.md b/docs/developer/adding-pgwire-engine.md new file mode 100644 index 000000000..cbc2309fc --- /dev/null +++ b/docs/developer/adding-pgwire-engine.md @@ -0,0 +1,109 @@ +# Adding A PostgreSQL Wire Engine + +PostgreSQL wire compatibility is a transport contract, not a SQL dialect or +session-state contract. Keep engine differences behind `pgwire.Engine` and +its optional interfaces. GreptimeDB in `pkg/backends/greptimedb` is the worked +example; `pkg/backends/postgres` remains the reference PostgreSQL engine. + +## Start With Relay + +Implement `Name`, `DefaultPort`, `Dialect`, `Defaults`, `TimeAxis` and +`TimeSemantics`. Return a nil `Analyzer` initially. Register the engine in +`pkg/backends/providers/registry` using the PostgreSQL native adapter. Prove +startup, authentication, simple/extended query forwarding, errors and connection +recovery against a real origin before enabling caching. + +Do not infer cancellation, transactions, type OIDs, precision or successful +setting changes from a PostgreSQL-shaped response. Observe the raw protocol +and compare with a direct connection. Capture only the relevant fields, not +passwords or reusable authentication/cancellation material. + +## Backend Model + +Providers that also serve HTTP implement `pgwire.HTTPEngine.SupportsHTTP`. +Their `origin_url` remains HTTP; `postgres.upstream_url` independently selects +the native endpoint. HTTP userinfo is not a source of native credentials. +Listener validation sets runtime HTTP/native mappings and validates only the +configured protocols, including terminal user-router targets. + +GreptimeDB also has a MySQL engine. Native registry lookup must therefore use +`GetForProvider(protocol, provider)` or enumerate `ForProvider(provider)`. +`GetByProvider` deliberately returns nil when a provider has several adapters; +never let map iteration choose which protocol to validate or run. + +Health probes must match the selected surface. Reuse the shared scheduler; +do not create a provider-specific polling loop. Include endpoint, engine, +authentication, TLS and result-affecting settings in native restart identity. + +## Session And Time Semantics + +`TimeAxis` maps only verified OIDs to timestamp/date/numeric time axes. +`TimeSemantics` describes engine guarantees, such as GreptimeDB's naive UTC +timestamps and lossless float text. Assumed settings are for nonconfigurable +guarantees, not convenient substitutes for unknown role or server defaults. + +Implement `SessionDefaultsEngine` when PostgreSQL's defaults query is not +valid upstream. GreptimeDB probes `SHOW TIMEZONE`, `SHOW DateStyle` and +`SHOW IntervalStyle`. Returned columns must match the declared names and +contain exactly one non-null row per result. + +Implement `SessionSettingsEngine` for different tracked settings, aliases, +neutral settings or `SET LOCAL` semantics. GreptimeDB's `time_zone` aliases +`timezone`; its local setting persists. Requested startup settings partition +the identity but are not accepted as effective values before observation. +Failed statements must not publish successful changes. Unknown state bypasses +caching conservatively. + +Implement `SessionAnalyzer.ForSession` if query eligibility depends on the +effective timezone. Keep it cheap and immutable. Test two sessions with +different settings, successful/failed SET, reset, startup overrides and +reconnects. GreptimeDB's pgwire microsecond text precision must not be reused +for its MySQL nanosecond text or HTTP timestamp representation. + +## Analyzer And Renderer + +Follow [SQL dialect adapters](sql-dialect-adapters.md). Keep parser AST types +out of the shared plan. Reuse Cockroach adapter bucket matchers, compact-duration +parsing and narrowly scoped post-render hooks when applicable. The exported +compact parser rejects signs, fractions, whitespace, zero components and +overflow; supply immutable dialect-specific fixed-duration units. + +Use clause rewriters only when the complete clause semantics are proved. +Unsupported `RANGE`, `ALIGN` and TQL shapes need no speculative delta plan. +Do not silently round HTTP partial buckets away merely because a dashboard's +pgwire mode intentionally consumes complete buckets. Preserve all non-time +predicates and result ordering. Rendering must be immutable and reentrant. + +## Tests And Corpus + +Add an origin target to `integration/pgwire_conformance_test.go`; declare +unsupported capabilities explicitly and prove them with a direct probe. +Assertions must compare direct and proxied values and check counters so a +passthrough response cannot masquerade as a cache hit. + +Put a versioned corpus in the backend's `testdata/compatibility` directory. +Use `pkg/testutil/sqlcompat.Run`, `CheckGrafanaMacros` and `Benchmark` with a +corpus path and a session-zone-aware analyzer. Each delta case records units, +cadence, phase, bounds and columns. The runner checks canonical identity and +renders and analyzes an exact cache extent. Capture actual Grafana +`executedQueryString` over an unaligned range; record plugin version, original +macros and unsupported cases instead of guessing the expansion. + +Run PostgreSQL regressions as well as new engine tests. Measure the shared +pgwire gate/hit path and new analyzer/render/model paths, excluding fixture +construction from timed loops. Run race tests, fuzz parsing boundaries, and +retain failed live attempts separately from corrected runs. + +## Developer Environment + +Use `-direct` and `-trickster` datasource names and dashboard +regex `/^-/`. Share the generated seed files and time-window metadata; +do not download another independent dataset. Give the seeder a writer role +and Grafana/Trickster read-only origin roles. + +Reserved ports are 8480 HTTP, 8485 Flight SQL, 8486 MySQL, 8487 ClickHouse, +8488 PostgreSQL/TimescaleDB, 8489 GreptimeDB pgwire and 8490 future QuestDB. +GreptimeDB MySQL uses 8491. Check host occupancy and bind remote acceptance +ports to loopback. Pin an upstream image that passes direct Grafana health +and queries before validating Trickster. Record nightly and stable builds +accurately; an upstream merge alone is not a released or tested binary. diff --git a/docs/developer/backend-extensibility.md b/docs/developer/backend-extensibility.md index a8f047499..91c233623 100644 --- a/docs/developer/backend-extensibility.md +++ b/docs/developer/backend-extensibility.md @@ -22,7 +22,7 @@ While this might sound daunting, it is actually much easier than it appears on t A Time Series Backend is used by Trickster to 1) manipulate HTTP requests and responses in order to accelerate the requests, 2) unmarshal data from backend databases into the [Common Time Series Format](https://github.com/trickstercache/trickster/blob/main/pkg/timeseries/dataset/dataset.go), and 3) marshal from the CTSF into a format supported by the Provider as requested by the downstream Client. -Trickster provides 1 required interfaces for enabling a new Provider: [Time Series Backend](https://github.com/trickstercache/trickster/blob/main/pkg/backends/timeseries_backend.go). Separately, you must implement `io.Writer`-based marshalers and unmarshalers that conform to Trickster's [Modeler specifications](https://github.com/trickstercache/trickster/blob/main/pkg/timeseries/modeler.go). +Trickster provides 1 required interfaces for enabling a new Provider: [Time Series Backend](https://github.com/trickstercache/trickster/blob/main/pkg/backends/timeseries_backend.go). Separately, you must implement `io.Writer`-based marshalers and unmarshalers that conform to Trickster's [Modeler specifications](https://github.com/trickstercache/trickster/blob/main/pkg/timeseries/modeler.go). The [Streaming Unmarshalers](./streaming-unmarshalers.md) guide covers the shared packages that decode a response body straight into a DataSet. Once data is unmarshaled into the Common Time Series Format, Trickster's other packages will handle operations like Delta Proxy Caching, etc. Thus, the implementer of a new Provider only needs to worry about wire protocols and formats. Specifically, you will need to know these things about the Backend: diff --git a/docs/developer/environment/README.md b/docs/developer/environment/README.md index b9fb98e9a..47c4baaa8 100644 --- a/docs/developer/environment/README.md +++ b/docs/developer/environment/README.md @@ -197,9 +197,13 @@ dev config registers a matching backend for each, so Grafana can query the upstream directly or via Trickster for a side-by-side comparison. +GreptimeDB also provides direct SQL and PromQL queries, a pgwire SQL cache, +and an HTTP proxy; +see [GreptimeDB Details](#greptimedb-details). + ## Seed data -ClickHouse, MySQL, TimescaleDB, and Druid are all loaded with the same +ClickHouse, MySQL, TimescaleDB, GreptimeDB, and Druid are all loaded with the same synthetic `trips` dataset: about 1.9 million cab rides in the fictional city of Emberwick over a 12-week window, with the same 45-column schema, label cardinality, and daily/weekly usage curve as a real ride dataset. Nothing is downloaded: the @@ -257,9 +261,9 @@ a container. See `hack/seedgen/README.md`. `make developer-seed-data` runs `hack/developer-seed-data.sh`, which regenerates the seed window and then runs every seeder concurrently. Set `SEED_TARGET` to a space- or comma-separated subset of `clickhouse`, `mysql`, -`timescaledb`, `druid`, `prometheus`, and `graphite` to scope the run, for example -`SEED_TARGET=timescaledb make developer-seed-data`. The seed instant is -recomputed on every run, so a scoped re-seed shifts only the selected +`timescaledb`, `greptimedb`, `druid`, `prometheus`, and `graphite` to scope the +run, for example `SEED_TARGET=timescaledb make developer-seed-data`. The seed +instant is recomputed on every run, so a scoped re-seed shifts only the selected databases; the others keep their previous shift, and dashboards that compare backends will disagree until a full run re-syncs them. @@ -329,7 +333,7 @@ Manager processes with Bash inside the one development container. The developer environment includes a pinned MySQL 8.4 (LTS) container seeded with the same auto-phased synthetic `trips` dataset used by ClickHouse, -TimescaleDB, and Druid. All four seeders read the shared generated files in +TimescaleDB, GreptimeDB, and Druid. All five seeders read the shared generated files in `docker-compose-data/seed-data`, so the data is generated once regardless of which seeder runs first (see [Seed data](#seed-data)). @@ -356,7 +360,7 @@ relationships in the relational copies while placing approximately half of the pickup distribution before and half after the seed instant. To re-seed (for example, after the data ages out of range), run `make developer-seed-data`, which first runs the `seed_data_generate` service and then reloads ClickHouse, -MySQL, TimescaleDB, and Druid in parallel. A Trickster started before the re-seed still +MySQL, TimescaleDB, GreptimeDB, and Druid in parallel. A Trickster started before the re-seed still holds the previous timeseries in its memory cache, so restart `make serve-dev` afterwards (or compare against a `-direct` datasource) to see the new data. @@ -529,6 +533,10 @@ the range is a `phit` for the part not held. The three object panels are a key, and a `hit` only on a fixed range. Saltmarrow is the sparsest borough in the synthetic data, so its 5-minute series has real gaps to fill. +Re-seeding waits for active startup seeders before regenerating their shared +fixture or reloading tables. A startup seeder failure stops that attempt; +inspect its logs before retrying. + Re-seeding rewrites the rows under whatever a running Trickster has cached. Restart Trickster after `make developer-seed-data`, or the direct and Trickster data sources will disagree over the ranges that were cached. @@ -674,3 +682,209 @@ is idempotent and always drops, re-creates, and reloads the `trips` table. PostgreSQL 18 images keep their data under `/var/lib/postgresql/18/docker`, so the `timescaledb-data` volume is mounted at `/var/lib/postgresql`; the init SQL only runs when that volume is empty (`make developer-delete` resets it). + +## GreptimeDB Details + +The developer environment pins the official +`greptime/greptimedb-nightly:nightly-20260923-e91faa9df` image by digest in +standalone mode, with UTC as the default time zone and telemetry disabled. +It contains the PostgreSQL health-query fix missing from v1.2.1. This is a +nightly build, not a stable release; update the pin only after repeating the +direct and proxied acceptance suites. + +| Endpoint | Address | +| --- | --- | +| HTTP SQL / health | `http://127.0.0.1:4000/v1/sql` / `/health` | +| Prometheus-compatible API | `http://127.0.0.1:4000/v1/prometheus` | +| gRPC | `127.0.0.1:4001` | +| MySQL wire | `127.0.0.1:4002` | +| PostgreSQL wire | `127.0.0.1:4003` | +| Trickster PostgreSQL wire | `127.0.0.1:8489` | +| Trickster MySQL wire | `127.0.0.1:8491` | +| Trickster HTTP | `http://127.0.0.1:8480/greptimedb1` | + +The database is `public`. The read-only mounted static user file contains +developer-only credentials: + +| User | Password | Access | +| --- | --- | --- | +| `seeder` | `trickster-dev-seed` | Schema creation and seeding | +| `grafana_ro` | `trickster-dev-grafana` | Read-only Grafana direct access | +| `trickster` | `trickster-dev-upstream` | Read-only Trickster upstream login | + +These are not production credentials. Data persists in the named +`greptimedb-data` volume; `make developer-delete` deletes it. Port 8489 is +Trickster's GreptimeDB pgwire listener. Port 8490 remains reserved for QuestDB. + +The `greptimedb1` backend serves the default HTTP listener and its own +pgwire and MySQL listeners. `origin_url` is the HTTP base; `postgres.upstream_url` supplies +the separate native host, credentials and database. Without that override, +pgwire uses a `postgres(ql)://` origin as before, or the HTTP origin's host +on port 4003 with no credentials. HTTP credentials, path and port are never +reused for pgwire. Terminated authentication and native health probes need +upstream credentials; passthrough sessions use the client's origin login. +The shared basic authenticator sets `proxy_preserve: true` so HTTP forwards +the client's valid origin credentials. It does not substitute pgwire's +upstream credentials into HTTP requests. + +`mysql.upstream_url` independently configures the MySQL origin, with port +4002 as its default. The MySQL listener terminates authentication and requires +an authenticator plus an explicit native upstream login. An HTTP origin's +credentials are never reused for either native protocol. MySQL-only deployments +can list just `greptimedb-mysql`; their health probe authenticates and sends +`COM_PING` to that native endpoint. + +MySQL SELECT queries use Vitess for analysis. The delta path supports UTC +`DATE_BIN('5m', ts, FROM_UNIXTIME(0))` and fixed second/minute/hour/day +`DATE_TRUNC` buckets with half-open `FROM_UNIXTIME(integer_seconds)` bounds, +rounded inward to complete buckets. Widths must be positive whole seconds. +Other deterministic SELECTs, including inclusive upper bounds, ranges with no +complete bucket and unverified timezones, use the object cache; +unknown functions and session state bypass caching. The adapter probes the +actual session timezone and preserves nine-digit timestamp text, large integers, +NULL ordering and bytewise string grouping. MySQL's `UNIX_TIMESTAMP` is not a +GreptimeDB function. The wire-compatible transaction commands do not establish +storage transactions; Trickster still bypasses caching while inside one. +As with the existing MySQL backend, prepared statements and multiple statements +are outside this listener's supported protocol contract. + +Run `sh hack/greptimedb-check.sh --mysql` against the isolated developer origin +for the native MySQL cache and session checks. The development read-only users +cannot `SET time_zone`; clients that issue it automatically receive the origin's +permission error. The adapter does not hide that error or increase their rights. + +The pgwire surface caches supported simple-protocol SQL queries. HTTP SQL +supports delta caching for typed `greptimedb_v1` responses and whole-response +caching for other eligible SELECT queries. Other HTTP surfaces remain +proxy-only. The backend preserves origin +paths such as `/v1/sql` and `/v1/prometheus/api/v1/query_range`; it adds no +bare `/api/v1` aliases. The mixed backend uses `GET /health`; a backend with +only a pgwire listener uses a native login probe instead. For pgwire-only +deployment, list only the postgres listener. For passthrough authentication, +omit `authenticator_name`; clients must then have valid origin credentials. + +Pgwire delta caching recognizes `date_bin(INTERVAL '5 minutes', col)`, compact +`date_bin('5m', col)`, UTC-session `date_trunc`, and Grafana's epoch-floor +expressions over timestamps or epoch-second columns. Fixed widths below one +microsecond and sub-microsecond bucket origins are object-only because pgwire +timestamp text does not preserve that precision. Variable calendar widths, +window functions and native `RANGE ... ALIGN` queries use the object cache; +`TQL`, writes and volatile expressions are not cached. As with PostgreSQL, +unaligned raw-time bounds are rounded inward to complete buckets. + +HTTP SQL uses the same complete-bucket alignment as pgwire. Ranges without a +complete bucket retain the original query. The delta model preserves typed columns, nulls, +integer precision and ordering, including empty results. Unsupported schemas +or a failed gap fetch retry the complete original SELECT instead of returning +an incomplete result. Database, timezone and authentication identity are +part of the cache key. Non-SELECT and multi-statement requests are not cached. +`format` values other than `greptimedb_v1` and requests with `limit` never use +delta caching. Authenticated GET object responses require origin permission +for shared caching; GreptimeDB's default responses do not grant it. + +PromQL uses `/v1/prometheus/api/v1/` and shares the Prometheus cache/model +implementation. Range endpoints round down to epoch-aligned steps while +retaining millisecond precision; +database URL/header selection and `lookback` are part of cache identity. URL +parameters take precedence over POST form parameters, and a form-only `db` +is ignored, matching GreptimeDB. Instant and metadata timestamps are not +rounded. Unknown parameters, unsupported methods, finer-than-millisecond grids +and `count_values` use passthrough. Aligned grids support configured time sharding; +fast-forward is disabled for fractional-second grids. + +An ALB using these paths can set `output_format: greptimedb` for TSM. Numeric +planning and reduction are shared with Prometheus, with Greptime-specific +parameter precedence and metric-name preservation. `count_values` cannot be +merged reliably while the origin omits its numeric grouping label. Finalizers +requiring dynamic metric-name discovery also reject the merge rather than +invent a combined result. A single backend still relays those requests. +Run `sh hack/greptimedb-check.sh --promql` for the isolated fixture-based +[PromQL acceptance suite](../../../integration/greptimedb/README.md#promql-provider-and-merge-acceptance). + +Terminated logins probe the actual timezone, date style and interval style using +`SHOW`, not PostgreSQL's `current_setting()`. A successful `SET`, including +GreptimeDB's persistent `SET LOCAL` and `time_zone` alias, updates the cache +identity. Failed changes do not. Non-ISO timestamp output falls back to the +object cache. Unknown session state disables caching. Passthrough sessions do +not run a hidden SQL probe; timezone-dependent analysis stays conservative +until the effective zone is known. + +```sh +curl --fail-with-body -u grafana_ro:trickster-dev-grafana \ + --data-urlencode 'sql=SELECT 1' http://127.0.0.1:4000/v1/sql +PGPASSWORD=trickster-dev-grafana psql -h 127.0.0.1 -p 4003 \ + -U grafana_ro -d public -c 'SELECT 1' +PGPASSWORD=trickster-dev-grafana psql -h 127.0.0.1 -p 8489 \ + -U grafana_ro -d public -c 'SELECT 1' +curl --fail-with-body -u grafana_ro:trickster-dev-grafana \ + --data-urlencode 'sql=SELECT 1' http://127.0.0.1:8480/greptimedb1/v1/sql +curl --fail-with-body -u grafana_ro:trickster-dev-grafana \ + --data-urlencode 'query=up' \ + http://127.0.0.1:8480/greptimedb1/v1/prometheus/api/v1/query +SEED_TARGET=greptimedb make developer-seed-data +``` + +`greptimedb_seed` runs the dependency-free `hack/greptimeseed` Go loader on +the same generated files and `SHIFT_SECONDS` as the other trips databases. +The loader uses HTTP SQL because v1.2.1 rejects line-protocol string fields +for existing `DATE` columns. It preserves the TimescaleDB column types, +uses microsecond timestamps and seconds for `pickup_epoch`, regenerates +dates in UTC, and retains empty text while mapping empty numeric fields +to NULL. `cab_type` and `vendor_id` are low-cardinality primary-key tags. +`append_mode='true'` keeps distinct trips at the same timestamp. + +Each seed drops and rebuilds `trips` and a static `trips_15m` rollup. It +checks exact row count, shifted pickup/dropoff bounds, centering on seed +time, date/epoch agreement and rollup coverage. INSERTs are not retried +after an ambiguous error, since append mode could duplicate a batch. Re-run +the seed to start from empty tables. Inspect failures with +`docker compose logs greptimedb greptimedb_seed`. + +Grafana provisions `greptimedb-direct` (`ds_greptimedb_direct`) using its +bundled PostgreSQL datasource, with the TimescaleDB option disabled, and +`greptimedb-prom-direct` (`ds_greptimedb_prom_direct`) using its bundled +Prometheus datasource. Both use the read-only account. Prometheus remote +write supplies the latter with samples from the existing scrape jobs. +No GreptimeDB-specific Grafana plugin is installed. + +Their proxy counterparts are `greptimedb-trickster` (`ds_greptimedb_trickster`) +and `greptimedb-prom-trickster` (`ds_greptimedb_prom_trickster`). The SQL pair +appears in the GreptimeDB dashboard selector without changing panel queries. + +The [GreptimeDB dashboard](http://127.0.0.1:3000/d/trickster-greptimedb) +preserves the TimescaleDB dashboard's panel IDs, layout and datasource +selector. Its SQL differences are: + +| Panel IDs | SQL adaptation | +| --- | --- | +| 1-4, 6 | Same Grafana macros and FILTER aggregates; disabling the TimescaleDB option expands time-group macros to epoch-floor expressions | +| 5 | `round(avg(...), 2)` without PostgreSQL's `::numeric` cast | +| 20-21 | Unchanged epoch-floor and `$__unixEpochGroupAlias` expressions | +| 22-23 | Native `RANGE '5m' FILL NULL/PREV ... ALIGN '5m'`, epoch-aligned with `BY ()` so tags do not split the result | +| 24 | `date_bin(INTERVAL '5 minutes', ...)` instead of `time_bucket`, retaining the window function | +| 25 | Static rollup, `date_bin(INTERVAL '15 minutes', ...)`, and quoted `"bucket"` (a GreptimeDB keyword) | + +Native RANGE filling covers the span of matching data, unlike TimescaleDB's +gapfill over the entire requested range; leading and trailing empty buckets +can therefore differ in panels 22-23. Performance panels use `greptimedb1` +and `greptimedb` labels. The native SQL cache populates the classification, +cache-status, request-duration and point counters. The rewrite-failure panel +has no series until a fallback is needed; that is not a missing-data failure. +Cache operations, utilization and evictions use the dedicated `greptimedb_fs` +filesystem cache. The memory provider does not export storage-usage gauges. +The eviction panel remains empty until an eviction occurs. + +### Direct Environment Checks + +Run `make developer-greptimedb-check` for repeatable, read-only validation of +the direct environment. The [acceptance guide](../../../integration/greptimedb/README.md) +describes its assertions, unique result directories and remaining human-review +checks. An empty datasource or a query error inside HTTP 200 fails the suite. + +GreptimeDB v1.2.1 rejects the comment-only `-- ping` query used by Grafana +13.1.3's PostgreSQL health check. The pinned official nightly contains +[GreptimeDB #9295](https://github.com/GreptimeTeam/greptimedb/pull/9295), which +repairs that response. The health check is tested directly against GreptimeDB, +without a Trickster shim. Retain old-version failures separately; an upstream +merge or a source-built repair is not evidence that a stable release contains +the fix. Changing the image requires repeating the full acceptance suite. diff --git a/docs/developer/environment/docker-compose-data/dashboards/trickster-greptimedb.json b/docs/developer/environment/docker-compose-data/dashboards/trickster-greptimedb.json new file mode 100644 index 000000000..5c44ddc3c --- /dev/null +++ b/docs/developer/environment/docker-compose-data/dashboards/trickster-greptimedb.json @@ -0,0 +1,1836 @@ +{ + "annotations": { + "list": [ + { + "builtIn": 1, + "datasource": { + "type": "grafana", + "uid": "-- Grafana --" + }, + "enable": true, + "hide": true, + "iconColor": "rgba(0, 211, 255, 1)", + "name": "Annotations & Alerts", + "type": "dashboard" + } + ] + }, + "description": "Synthetic Emberwick trips seeded into GreptimeDB with a source-derived shift centered on seed time. Set an absolute future range to inspect the future half.", + "editable": true, + "fiscalYearStartMonth": 0, + "graphTooltip": 0, + "links": [], + "panels": [ + { + "collapsed": false, + "gridPos": { + "h": 1, + "w": 24, + "x": 0, + "y": 0 + }, + "id": 100, + "title": "Trips Data", + "type": "row" + }, + { + "datasource": { + "type": "grafana-postgresql-datasource", + "uid": "${datasource}" + }, + "description": "Fixed 5-minute buckets with a half-open time range, the delta-cache friendly pattern.", + "fieldConfig": { + "defaults": { + "color": { + "mode": "palette-classic" + }, + "mappings": [], + "thresholds": { + "mode": "absolute", + "steps": [ + { + "color": "green" + }, + { + "color": "red", + "value": 80 + } + ] + }, + "custom": { + "axisBorderShow": false, + "axisCenteredZero": false, + "axisColorMode": "text", + "axisLabel": "", + "axisPlacement": "auto", + "barAlignment": 0, + "barWidthFactor": 0.6, + "drawStyle": "line", + "fillOpacity": 0, + "gradientMode": "none", + "hideFrom": { + "legend": false, + "tooltip": false, + "viz": false + }, + "insertNulls": false, + "lineInterpolation": "linear", + "lineWidth": 1, + "pointSize": 5, + "scaleDistribution": { + "type": "linear" + }, + "showPoints": "auto", + "spanNulls": false, + "stacking": { + "group": "A", + "mode": "none" + }, + "thresholdsStyle": { + "mode": "off" + } + } + }, + "overrides": [] + }, + "gridPos": { + "h": 10, + "w": 12, + "x": 0, + "y": 1 + }, + "id": 1, + "options": { + "legend": { + "calcs": [], + "displayMode": "list", + "placement": "bottom", + "showLegend": true + }, + "tooltip": { + "hideZeros": false, + "mode": "multi", + "sort": "none" + } + }, + "targets": [ + { + "datasource": { + "type": "grafana-postgresql-datasource", + "uid": "${datasource}" + }, + "editorMode": "code", + "format": "time_series", + "rawQuery": true, + "rawSql": "SELECT\n $__timeGroupAlias(pickup_datetime, '5m'),\n count(*) AS trips\nFROM trips\nWHERE pickup_datetime >= $__timeFrom() AND pickup_datetime < $__timeTo()\nGROUP BY time\nORDER BY time", + "refId": "A", + "sql": { + "columns": [ + { + "parameters": [], + "type": "function" + } + ], + "groupBy": [ + { + "property": { + "type": "string" + }, + "type": "groupBy" + } + ], + "limit": 50 + } + } + ], + "title": "# Trips Over Time (5m buckets)", + "type": "timeseries" + }, + { + "datasource": { + "type": "grafana-postgresql-datasource", + "uid": "${datasource}" + }, + "description": "Grafana-driven dynamic bucket size via $__interval, with a 1m minimum interval on the panel.", + "fieldConfig": { + "defaults": { + "color": { + "mode": "palette-classic" + }, + "mappings": [], + "thresholds": { + "mode": "absolute", + "steps": [ + { + "color": "green" + }, + { + "color": "red", + "value": 80 + } + ] + }, + "custom": { + "axisBorderShow": false, + "axisCenteredZero": false, + "axisColorMode": "text", + "axisLabel": "", + "axisPlacement": "auto", + "barAlignment": 0, + "barWidthFactor": 0.6, + "drawStyle": "line", + "fillOpacity": 0, + "gradientMode": "none", + "hideFrom": { + "legend": false, + "tooltip": false, + "viz": false + }, + "insertNulls": false, + "lineInterpolation": "linear", + "lineWidth": 1, + "pointSize": 5, + "scaleDistribution": { + "type": "linear" + }, + "showPoints": "auto", + "spanNulls": false, + "stacking": { + "group": "A", + "mode": "none" + }, + "thresholdsStyle": { + "mode": "off" + } + } + }, + "overrides": [] + }, + "gridPos": { + "h": 10, + "w": 12, + "x": 12, + "y": 1 + }, + "id": 2, + "options": { + "legend": { + "calcs": [], + "displayMode": "list", + "placement": "bottom", + "showLegend": true + }, + "tooltip": { + "hideZeros": false, + "mode": "multi", + "sort": "none" + } + }, + "targets": [ + { + "datasource": { + "type": "grafana-postgresql-datasource", + "uid": "${datasource}" + }, + "editorMode": "code", + "format": "time_series", + "rawQuery": true, + "rawSql": "SELECT\n $__timeGroupAlias(pickup_datetime, $__interval),\n count(*) AS trips\nFROM trips\nWHERE pickup_datetime >= $__timeFrom() AND pickup_datetime < $__timeTo()\nGROUP BY time\nORDER BY time", + "refId": "A", + "sql": { + "columns": [ + { + "parameters": [], + "type": "function" + } + ], + "groupBy": [ + { + "property": { + "type": "string" + }, + "type": "groupBy" + } + ], + "limit": 50 + } + } + ], + "title": "# Trips Over Time ($__interval)", + "type": "timeseries", + "interval": "1m" + }, + { + "datasource": { + "type": "grafana-postgresql-datasource", + "uid": "${datasource}" + }, + "description": "Multiple numeric value columns in a single time-series result.", + "fieldConfig": { + "defaults": { + "color": { + "mode": "palette-classic" + }, + "mappings": [], + "thresholds": { + "mode": "absolute", + "steps": [ + { + "color": "green" + }, + { + "color": "red", + "value": 80 + } + ] + }, + "custom": { + "axisBorderShow": false, + "axisCenteredZero": false, + "axisColorMode": "text", + "axisLabel": "", + "axisPlacement": "auto", + "barAlignment": 0, + "barWidthFactor": 0.6, + "drawStyle": "line", + "fillOpacity": 0, + "gradientMode": "none", + "hideFrom": { + "legend": false, + "tooltip": false, + "viz": false + }, + "insertNulls": false, + "lineInterpolation": "linear", + "lineWidth": 1, + "pointSize": 5, + "scaleDistribution": { + "type": "linear" + }, + "showPoints": "auto", + "spanNulls": false, + "stacking": { + "group": "A", + "mode": "none" + }, + "thresholdsStyle": { + "mode": "off" + } + }, + "unit": "currencyUSD" + }, + "overrides": [] + }, + "gridPos": { + "h": 10, + "w": 12, + "x": 0, + "y": 11 + }, + "id": 3, + "options": { + "legend": { + "calcs": [], + "displayMode": "list", + "placement": "bottom", + "showLegend": true + }, + "tooltip": { + "hideZeros": false, + "mode": "multi", + "sort": "none" + } + }, + "targets": [ + { + "datasource": { + "type": "grafana-postgresql-datasource", + "uid": "${datasource}" + }, + "editorMode": "code", + "format": "time_series", + "rawQuery": true, + "rawSql": "SELECT\n $__timeGroupAlias(pickup_datetime, '15m'),\n avg(fare_amount) AS avg_fare,\n avg(tip_amount) AS avg_tip,\n avg(total_amount) AS avg_total\nFROM trips\nWHERE pickup_datetime >= $__timeFrom() AND pickup_datetime < $__timeTo()\nGROUP BY time\nORDER BY time", + "refId": "A", + "sql": { + "columns": [ + { + "parameters": [], + "type": "function" + } + ], + "groupBy": [ + { + "property": { + "type": "string" + }, + "type": "groupBy" + } + ], + "limit": 50 + } + } + ], + "title": "Fare Metrics (15m buckets)", + "type": "timeseries" + }, + { + "datasource": { + "type": "grafana-postgresql-datasource", + "uid": "${datasource}" + }, + "description": "A string dimension (cab_type) grouped with time to produce multiple series from one query.", + "fieldConfig": { + "defaults": { + "color": { + "mode": "palette-classic" + }, + "mappings": [], + "thresholds": { + "mode": "absolute", + "steps": [ + { + "color": "green" + }, + { + "color": "red", + "value": 80 + } + ] + }, + "custom": { + "axisBorderShow": false, + "axisCenteredZero": false, + "axisColorMode": "text", + "axisLabel": "", + "axisPlacement": "auto", + "barAlignment": 0, + "barWidthFactor": 0.6, + "drawStyle": "line", + "fillOpacity": 0, + "gradientMode": "none", + "hideFrom": { + "legend": false, + "tooltip": false, + "viz": false + }, + "insertNulls": false, + "lineInterpolation": "linear", + "lineWidth": 1, + "pointSize": 5, + "scaleDistribution": { + "type": "linear" + }, + "showPoints": "auto", + "spanNulls": false, + "stacking": { + "group": "A", + "mode": "none" + }, + "thresholdsStyle": { + "mode": "off" + } + } + }, + "overrides": [ + { + "matcher": { + "id": "byRegexp", + "options": "/orange/" + }, + "properties": [ + { + "id": "color", + "value": { + "fixedColor": "orange", + "mode": "fixed" + } + } + ] + }, + { + "matcher": { + "id": "byRegexp", + "options": "/blue/" + }, + "properties": [ + { + "id": "color", + "value": { + "fixedColor": "blue", + "mode": "fixed" + } + } + ] + }, + { + "matcher": { + "id": "byRegexp", + "options": "/purple/" + }, + "properties": [ + { + "id": "color", + "value": { + "fixedColor": "purple", + "mode": "fixed" + } + } + ] + } + ] + }, + "gridPos": { + "h": 10, + "w": 12, + "x": 12, + "y": 11 + }, + "id": 4, + "options": { + "legend": { + "calcs": [], + "displayMode": "list", + "placement": "bottom", + "showLegend": true + }, + "tooltip": { + "hideZeros": false, + "mode": "multi", + "sort": "none" + } + }, + "targets": [ + { + "datasource": { + "type": "grafana-postgresql-datasource", + "uid": "${datasource}" + }, + "editorMode": "code", + "format": "time_series", + "rawQuery": true, + "rawSql": "SELECT\n $__timeGroupAlias(pickup_datetime, '15m'),\n cab_type AS metric,\n count(*) AS trips\nFROM trips\nWHERE pickup_datetime >= $__timeFrom() AND pickup_datetime < $__timeTo()\nGROUP BY time, cab_type\nORDER BY time", + "refId": "A", + "sql": { + "columns": [ + { + "parameters": [], + "type": "function" + } + ], + "groupBy": [ + { + "property": { + "type": "string" + }, + "type": "groupBy" + } + ], + "limit": 50 + } + } + ], + "title": "Trips by Cab Type (15m buckets)", + "type": "timeseries" + }, + { + "datasource": { + "type": "grafana-postgresql-datasource", + "uid": "${datasource}" + }, + "description": "Table-format query over the dashboard time range; not a time series.", + "fieldConfig": { + "defaults": { + "color": { + "mode": "palette-classic" + }, + "mappings": [], + "thresholds": { + "mode": "absolute", + "steps": [ + { + "color": "green" + }, + { + "color": "red", + "value": 80 + } + ] + } + }, + "overrides": [] + }, + "gridPos": { + "h": 10, + "w": 24, + "x": 0, + "y": 21 + }, + "id": 5, + "options": { + "cellHeight": "sm", + "footer": { + "countRows": false, + "fields": "", + "reducer": [ + "sum" + ], + "show": false + }, + "showHeader": true + }, + "targets": [ + { + "datasource": { + "type": "grafana-postgresql-datasource", + "uid": "${datasource}" + }, + "editorMode": "code", + "format": "table", + "rawQuery": true, + "rawSql": "SELECT\n pickup_neighborhood_name AS neighborhood,\n count(*) AS trips,\n round(avg(trip_distance), 2) AS avg_distance_mi,\n round(avg(total_amount), 2) AS avg_total_usd\nFROM trips\nWHERE $__timeFilter(pickup_datetime)\nGROUP BY pickup_neighborhood_name\nORDER BY trips DESC\nLIMIT 10", + "refId": "A", + "sql": { + "columns": [ + { + "parameters": [], + "type": "function" + } + ], + "groupBy": [ + { + "property": { + "type": "string" + }, + "type": "groupBy" + } + ], + "limit": 50 + } + } + ], + "title": "Top Pickup Neighborhoods", + "type": "table" + }, + { + "datasource": { + "type": "grafana-postgresql-datasource", + "uid": "${datasource}" + }, + "description": "Share of trips per bucket not coded as cash; a conditional aggregate in one query.", + "fieldConfig": { + "defaults": { + "color": { + "mode": "palette-classic" + }, + "mappings": [], + "thresholds": { + "mode": "absolute", + "steps": [ + { + "color": "green" + }, + { + "color": "red", + "value": 80 + } + ] + }, + "custom": { + "axisBorderShow": false, + "axisCenteredZero": false, + "axisColorMode": "text", + "axisLabel": "", + "axisPlacement": "auto", + "barAlignment": 0, + "barWidthFactor": 0.6, + "drawStyle": "line", + "fillOpacity": 0, + "gradientMode": "none", + "hideFrom": { + "legend": false, + "tooltip": false, + "viz": false + }, + "insertNulls": false, + "lineInterpolation": "linear", + "lineWidth": 1, + "pointSize": 5, + "scaleDistribution": { + "type": "linear" + }, + "showPoints": "auto", + "spanNulls": false, + "stacking": { + "group": "A", + "mode": "none" + }, + "thresholdsStyle": { + "mode": "off" + } + }, + "unit": "percent", + "max": 100, + "min": 0 + }, + "overrides": [] + }, + "gridPos": { + "h": 10, + "w": 24, + "x": 0, + "y": 31 + }, + "id": 6, + "options": { + "legend": { + "calcs": [], + "displayMode": "list", + "placement": "bottom", + "showLegend": true + }, + "tooltip": { + "hideZeros": false, + "mode": "multi", + "sort": "none" + } + }, + "targets": [ + { + "datasource": { + "type": "grafana-postgresql-datasource", + "uid": "${datasource}" + }, + "editorMode": "code", + "format": "time_series", + "rawQuery": true, + "rawSql": "SELECT\n $__timeGroupAlias(pickup_datetime, '5m'),\n 100.0 * count(*) FILTER (WHERE payment_type <> 'CSH') / count(*) AS card_use_rate\nFROM trips\nWHERE pickup_datetime >= $__timeFrom() AND pickup_datetime < $__timeTo()\nGROUP BY time\nORDER BY time", + "refId": "A", + "sql": { + "columns": [ + { + "parameters": [], + "type": "function" + } + ], + "groupBy": [ + { + "property": { + "type": "string" + }, + "type": "groupBy" + } + ], + "limit": 50 + } + } + ], + "title": "% Trips Paid w/ Credit Card (5m buckets)", + "type": "timeseries" + }, + { + "collapsed": false, + "gridPos": { + "h": 1, + "w": 24, + "x": 0, + "y": 41 + }, + "id": 102, + "title": "Cache Path Exercises", + "type": "row" + }, + { + "datasource": { + "type": "grafana-postgresql-datasource", + "uid": "${datasource}" + }, + "description": "Expected: delta. The bucket Grafana's $__timeGroup writes with the TimescaleDB option off: a number of epoch seconds rather than a timestamp. $__timeFilter expands to BETWEEN; Trickster normalizes the range to complete buckets and renders the inclusive upper bound below the next bucket.", + "fieldConfig": { + "defaults": { + "color": { + "mode": "palette-classic" + }, + "mappings": [], + "thresholds": { + "mode": "absolute", + "steps": [ + { + "color": "green" + }, + { + "color": "red", + "value": 80 + } + ] + }, + "custom": { + "axisBorderShow": false, + "axisCenteredZero": false, + "axisColorMode": "text", + "axisLabel": "", + "axisPlacement": "auto", + "barAlignment": 0, + "barWidthFactor": 0.6, + "drawStyle": "line", + "fillOpacity": 0, + "gradientMode": "none", + "hideFrom": { + "legend": false, + "tooltip": false, + "viz": false + }, + "insertNulls": false, + "lineInterpolation": "linear", + "lineWidth": 1, + "pointSize": 5, + "scaleDistribution": { + "type": "linear" + }, + "showPoints": "auto", + "spanNulls": false, + "stacking": { + "group": "A", + "mode": "none" + }, + "thresholdsStyle": { + "mode": "off" + } + } + }, + "overrides": [] + }, + "gridPos": { + "h": 10, + "w": 12, + "x": 0, + "y": 42 + }, + "id": 20, + "options": { + "legend": { + "calcs": [], + "displayMode": "list", + "placement": "bottom", + "showLegend": true + }, + "tooltip": { + "hideZeros": false, + "mode": "multi", + "sort": "none" + } + }, + "targets": [ + { + "datasource": { + "type": "grafana-postgresql-datasource", + "uid": "${datasource}" + }, + "editorMode": "code", + "format": "time_series", + "rawQuery": true, + "rawSql": "SELECT\n floor(extract(epoch from pickup_datetime)/300)*300 AS \"time\",\n count(*) AS trips\nFROM trips\nWHERE $__timeFilter(pickup_datetime)\nGROUP BY 1\nORDER BY 1", + "refId": "A", + "sql": { + "columns": [ + { + "parameters": [], + "type": "function" + } + ], + "groupBy": [ + { + "property": { + "type": "string" + }, + "type": "groupBy" + } + ], + "limit": 50 + } + } + ], + "title": "Epoch Floor Buckets ($__timeFilter)", + "type": "timeseries" + }, + { + "datasource": { + "type": "grafana-postgresql-datasource", + "uid": "${datasource}" + }, + "description": "Expected: delta. An integer epoch-seconds time column with $__unixEpochFilter bounds; the bucket is floor((col)/600)*600.", + "fieldConfig": { + "defaults": { + "color": { + "mode": "palette-classic" + }, + "mappings": [], + "thresholds": { + "mode": "absolute", + "steps": [ + { + "color": "green" + }, + { + "color": "red", + "value": 80 + } + ] + }, + "custom": { + "axisBorderShow": false, + "axisCenteredZero": false, + "axisColorMode": "text", + "axisLabel": "", + "axisPlacement": "auto", + "barAlignment": 0, + "barWidthFactor": 0.6, + "drawStyle": "line", + "fillOpacity": 0, + "gradientMode": "none", + "hideFrom": { + "legend": false, + "tooltip": false, + "viz": false + }, + "insertNulls": false, + "lineInterpolation": "linear", + "lineWidth": 1, + "pointSize": 5, + "scaleDistribution": { + "type": "linear" + }, + "showPoints": "auto", + "spanNulls": false, + "stacking": { + "group": "A", + "mode": "none" + }, + "thresholdsStyle": { + "mode": "off" + } + } + }, + "overrides": [] + }, + "gridPos": { + "h": 10, + "w": 12, + "x": 12, + "y": 42 + }, + "id": 21, + "options": { + "legend": { + "calcs": [], + "displayMode": "list", + "placement": "bottom", + "showLegend": true + }, + "tooltip": { + "hideZeros": false, + "mode": "multi", + "sort": "none" + } + }, + "targets": [ + { + "datasource": { + "type": "grafana-postgresql-datasource", + "uid": "${datasource}" + }, + "editorMode": "code", + "format": "time_series", + "rawQuery": true, + "rawSql": "SELECT\n $__unixEpochGroupAlias(pickup_epoch, '10m'),\n avg(total_amount) AS avg_total\nFROM trips\nWHERE $__unixEpochFilter(pickup_epoch)\nGROUP BY 1\nORDER BY 1", + "refId": "A", + "sql": { + "columns": [ + { + "parameters": [], + "type": "function" + } + ], + "groupBy": [ + { + "property": { + "type": "string" + }, + "type": "groupBy" + } + ], + "limit": 50 + } + } + ], + "title": "Avg Total ($__unixEpochGroup, 10m buckets)", + "type": "timeseries" + }, + { + "datasource": { + "type": "grafana-postgresql-datasource", + "uid": "${datasource}" + }, + "description": "Expected: object / invalid_sql. GreptimeDB RANGE ... ALIGN with FILL NULL fills gaps within the matching data span. Unlike TimescaleDB gapfill, it may omit leading and trailing empty buckets. Native RANGE clause rewriting is not enabled.", + "fieldConfig": { + "defaults": { + "color": { + "mode": "palette-classic" + }, + "mappings": [], + "thresholds": { + "mode": "absolute", + "steps": [ + { + "color": "green" + }, + { + "color": "red", + "value": 80 + } + ] + }, + "custom": { + "axisBorderShow": false, + "axisCenteredZero": false, + "axisColorMode": "text", + "axisLabel": "", + "axisPlacement": "auto", + "barAlignment": 0, + "barWidthFactor": 0.6, + "drawStyle": "line", + "fillOpacity": 0, + "gradientMode": "none", + "hideFrom": { + "legend": false, + "tooltip": false, + "viz": false + }, + "insertNulls": false, + "lineInterpolation": "linear", + "lineWidth": 1, + "pointSize": 5, + "scaleDistribution": { + "type": "linear" + }, + "showPoints": "auto", + "spanNulls": false, + "stacking": { + "group": "A", + "mode": "none" + }, + "thresholdsStyle": { + "mode": "off" + } + } + }, + "overrides": [] + }, + "gridPos": { + "h": 10, + "w": 12, + "x": 0, + "y": 52 + }, + "id": 22, + "options": { + "legend": { + "calcs": [], + "displayMode": "list", + "placement": "bottom", + "showLegend": true + }, + "tooltip": { + "hideZeros": false, + "mode": "multi", + "sort": "none" + } + }, + "targets": [ + { + "datasource": { + "type": "grafana-postgresql-datasource", + "uid": "${datasource}" + }, + "editorMode": "code", + "format": "time_series", + "rawQuery": true, + "rawSql": "SELECT\n pickup_datetime AS \"time\",\n count(*) RANGE '5m' FILL NULL AS trips\nFROM trips\nWHERE pickup_datetime >= $__timeFrom() AND pickup_datetime < $__timeTo()\n AND pickup_borough_name = 'Saltmarrow'\nALIGN '5m' TO '1970-01-01T00:00:00Z' BY ()\nORDER BY 1", + "refId": "A", + "sql": { + "columns": [ + { + "parameters": [], + "type": "function" + } + ], + "groupBy": [ + { + "property": { + "type": "string" + }, + "type": "groupBy" + } + ], + "limit": 50 + } + } + ], + "title": "Saltmarrow Trips, Gap-Filled (5m buckets)", + "type": "timeseries" + }, + { + "datasource": { + "type": "grafana-postgresql-datasource", + "uid": "${datasource}" + }, + "description": "Expected: object / invalid_sql. GreptimeDB RANGE ... ALIGN with FILL PREV carries the preceding value across gaps within the matching data span. A sub-range can lack that preceding value, so this query is cached whole.", + "fieldConfig": { + "defaults": { + "color": { + "mode": "palette-classic" + }, + "mappings": [], + "thresholds": { + "mode": "absolute", + "steps": [ + { + "color": "green" + }, + { + "color": "red", + "value": 80 + } + ] + }, + "custom": { + "axisBorderShow": false, + "axisCenteredZero": false, + "axisColorMode": "text", + "axisLabel": "", + "axisPlacement": "auto", + "barAlignment": 0, + "barWidthFactor": 0.6, + "drawStyle": "line", + "fillOpacity": 0, + "gradientMode": "none", + "hideFrom": { + "legend": false, + "tooltip": false, + "viz": false + }, + "insertNulls": false, + "lineInterpolation": "linear", + "lineWidth": 1, + "pointSize": 5, + "scaleDistribution": { + "type": "linear" + }, + "showPoints": "auto", + "spanNulls": false, + "stacking": { + "group": "A", + "mode": "none" + }, + "thresholdsStyle": { + "mode": "off" + } + } + }, + "overrides": [] + }, + "gridPos": { + "h": 10, + "w": 12, + "x": 12, + "y": 52 + }, + "id": 23, + "options": { + "legend": { + "calcs": [], + "displayMode": "list", + "placement": "bottom", + "showLegend": true + }, + "tooltip": { + "hideZeros": false, + "mode": "multi", + "sort": "none" + } + }, + "targets": [ + { + "datasource": { + "type": "grafana-postgresql-datasource", + "uid": "${datasource}" + }, + "editorMode": "code", + "format": "time_series", + "rawQuery": true, + "rawSql": "SELECT\n pickup_datetime AS \"time\",\n count(*) RANGE '5m' FILL PREV AS trips\nFROM trips\nWHERE pickup_datetime >= $__timeFrom() AND pickup_datetime < $__timeTo()\n AND pickup_borough_name = 'Saltmarrow'\nALIGN '5m' TO '1970-01-01T00:00:00Z' BY ()\nORDER BY 1", + "refId": "A", + "sql": { + "columns": [ + { + "parameters": [], + "type": "function" + } + ], + "groupBy": [ + { + "property": { + "type": "string" + }, + "type": "groupBy" + } + ], + "limit": 50 + } + } + ], + "title": "Saltmarrow Trips, Last Value Carried Forward", + "type": "timeseries" + }, + { + "datasource": { + "type": "grafana-postgresql-datasource", + "uid": "${datasource}" + }, + "description": "Expected: object / unsupported_format. A window frame spans buckets, so per-extent fetches would compute different values at the seams.", + "fieldConfig": { + "defaults": { + "color": { + "mode": "palette-classic" + }, + "mappings": [], + "thresholds": { + "mode": "absolute", + "steps": [ + { + "color": "green" + }, + { + "color": "red", + "value": 80 + } + ] + }, + "custom": { + "axisBorderShow": false, + "axisCenteredZero": false, + "axisColorMode": "text", + "axisLabel": "", + "axisPlacement": "auto", + "barAlignment": 0, + "barWidthFactor": 0.6, + "drawStyle": "line", + "fillOpacity": 0, + "gradientMode": "none", + "hideFrom": { + "legend": false, + "tooltip": false, + "viz": false + }, + "insertNulls": false, + "lineInterpolation": "linear", + "lineWidth": 1, + "pointSize": 5, + "scaleDistribution": { + "type": "linear" + }, + "showPoints": "auto", + "spanNulls": false, + "stacking": { + "group": "A", + "mode": "none" + }, + "thresholdsStyle": { + "mode": "off" + } + } + }, + "overrides": [] + }, + "gridPos": { + "h": 10, + "w": 12, + "x": 0, + "y": 62 + }, + "id": 24, + "options": { + "legend": { + "calcs": [], + "displayMode": "list", + "placement": "bottom", + "showLegend": true + }, + "tooltip": { + "hideZeros": false, + "mode": "multi", + "sort": "none" + } + }, + "targets": [ + { + "datasource": { + "type": "grafana-postgresql-datasource", + "uid": "${datasource}" + }, + "editorMode": "code", + "format": "time_series", + "rawQuery": true, + "rawSql": "SELECT\n date_bin(INTERVAL '5 minutes', pickup_datetime) AS \"time\",\n avg(count(*)) OVER (ORDER BY date_bin(INTERVAL '5 minutes', pickup_datetime) ROWS BETWEEN 11 PRECEDING AND CURRENT ROW) AS trips_1h_avg\nFROM trips\nWHERE pickup_datetime >= $__timeFrom() AND pickup_datetime < $__timeTo()\nGROUP BY 1\nORDER BY 1", + "refId": "A", + "sql": { + "columns": [ + { + "parameters": [], + "type": "function" + } + ], + "groupBy": [ + { + "property": { + "type": "string" + }, + "type": "groupBy" + } + ], + "limit": 50 + } + } + ], + "title": "Trips, 1h Moving Average (window function)", + "type": "timeseries" + }, + { + "datasource": { + "type": "grafana-postgresql-datasource", + "uid": "${datasource}" + }, + "description": "Expected: delta. The same series as 'Trips by Cab Type (15m buckets)', read from the static trips_15m rollup populated during seeding. The rollup is not continuously refreshed; the native SQL cache preserves the quoted bucket identifier.", + "fieldConfig": { + "defaults": { + "color": { + "mode": "palette-classic" + }, + "mappings": [], + "thresholds": { + "mode": "absolute", + "steps": [ + { + "color": "green" + }, + { + "color": "red", + "value": 80 + } + ] + }, + "custom": { + "axisBorderShow": false, + "axisCenteredZero": false, + "axisColorMode": "text", + "axisLabel": "", + "axisPlacement": "auto", + "barAlignment": 0, + "barWidthFactor": 0.6, + "drawStyle": "line", + "fillOpacity": 0, + "gradientMode": "none", + "hideFrom": { + "legend": false, + "tooltip": false, + "viz": false + }, + "insertNulls": false, + "lineInterpolation": "linear", + "lineWidth": 1, + "pointSize": 5, + "scaleDistribution": { + "type": "linear" + }, + "showPoints": "auto", + "spanNulls": false, + "stacking": { + "group": "A", + "mode": "none" + }, + "thresholdsStyle": { + "mode": "off" + } + } + }, + "overrides": [] + }, + "gridPos": { + "h": 10, + "w": 12, + "x": 12, + "y": 62 + }, + "id": 25, + "options": { + "legend": { + "calcs": [], + "displayMode": "list", + "placement": "bottom", + "showLegend": true + }, + "tooltip": { + "hideZeros": false, + "mode": "multi", + "sort": "none" + } + }, + "targets": [ + { + "datasource": { + "type": "grafana-postgresql-datasource", + "uid": "${datasource}" + }, + "editorMode": "code", + "format": "time_series", + "rawQuery": true, + "rawSql": "SELECT\n date_bin(INTERVAL '15 minutes', \"bucket\") AS \"time\",\n cab_type AS metric,\n sum(trips) AS trips\nFROM trips_15m\nWHERE \"bucket\" >= $__timeFrom() AND \"bucket\" < $__timeTo()\nGROUP BY 1, 2\nORDER BY 1", + "refId": "A", + "sql": { + "columns": [ + { + "parameters": [], + "type": "function" + } + ], + "groupBy": [ + { + "property": { + "type": "string" + }, + "type": "groupBy" + } + ], + "limit": 50 + } + } + ], + "title": "Trips by Cab Type (trips_15m rollup)", + "type": "timeseries" + }, + { + "collapsed": false, + "gridPos": { + "h": 1, + "w": 24, + "x": 0, + "y": 72 + }, + "id": 101, + "title": "Trickster Performance", + "type": "row" + }, + { + "datasource": { + "type": "prometheus", + "uid": "ds_prom_direct" + }, + "fieldConfig": { + "defaults": { + "unit": "ops" + }, + "overrides": [] + }, + "gridPos": { + "h": 9, + "w": 12, + "x": 0, + "y": 73 + }, + "id": 7, + "options": { + "legend": { + "calcs": [], + "displayMode": "list", + "placement": "bottom", + "showLegend": true + }, + "tooltip": { + "mode": "multi", + "sort": "none" + } + }, + "targets": [ + { + "datasource": { + "type": "prometheus", + "uid": "ds_prom_direct" + }, + "editorMode": "code", + "expr": "sum by (cache_mode, reason) (rate(trickster_sql_query_analysis_total{backend_name=\"greptimedb1\",dialect=\"greptimedb\"}[$__rate_interval]))", + "legendFormat": "{{cache_mode}} / {{reason}}", + "range": true, + "refId": "A" + } + ], + "title": "SQL Analysis Classifications", + "type": "timeseries" + }, + { + "datasource": { + "type": "prometheus", + "uid": "ds_prom_direct" + }, + "fieldConfig": { + "defaults": { + "unit": "ops" + }, + "overrides": [] + }, + "gridPos": { + "h": 9, + "w": 12, + "x": 12, + "y": 73 + }, + "id": 8, + "options": { + "legend": { + "calcs": [], + "displayMode": "list", + "placement": "bottom", + "showLegend": true + }, + "tooltip": { + "mode": "multi", + "sort": "none" + } + }, + "targets": [ + { + "datasource": { + "type": "prometheus", + "uid": "ds_prom_direct" + }, + "editorMode": "code", + "expr": "sum by (reason) (rate(trickster_sql_query_rewrite_failures_total{backend_name=\"greptimedb1\",dialect=\"greptimedb\"}[$__rate_interval]))", + "legendFormat": "{{reason}}", + "range": true, + "refId": "A" + } + ], + "title": "SQL Rewrite Failures", + "type": "timeseries" + }, + { + "datasource": { + "type": "prometheus", + "uid": "ds_prom_direct" + }, + "fieldConfig": { + "defaults": { + "unit": "ops" + }, + "overrides": [] + }, + "gridPos": { + "h": 9, + "w": 12, + "x": 0, + "y": 82 + }, + "id": 9, + "options": { + "legend": { + "calcs": [], + "displayMode": "list", + "placement": "bottom", + "showLegend": true + }, + "tooltip": { + "mode": "multi", + "sort": "none" + } + }, + "targets": [ + { + "datasource": { + "type": "prometheus", + "uid": "ds_prom_direct" + }, + "editorMode": "code", + "expr": "sum by (cache_mode, cache_status) (rate(trickster_sql_query_cache_total{backend_name=\"greptimedb1\",dialect=\"greptimedb\"}[$__rate_interval]))", + "legendFormat": "{{cache_mode}} / {{cache_status}}", + "range": true, + "refId": "A" + } + ], + "title": "OPC / DPC Cache Status", + "type": "timeseries" + }, + { + "datasource": { + "type": "prometheus", + "uid": "ds_prom_direct" + }, + "fieldConfig": { + "defaults": { + "unit": "s" + }, + "overrides": [] + }, + "gridPos": { + "h": 9, + "w": 12, + "x": 12, + "y": 82 + }, + "id": 10, + "options": { + "legend": { + "calcs": [], + "displayMode": "list", + "placement": "bottom", + "showLegend": true + }, + "tooltip": { + "mode": "multi", + "sort": "none" + } + }, + "targets": [ + { + "datasource": { + "type": "prometheus", + "uid": "ds_prom_direct" + }, + "editorMode": "code", + "expr": "histogram_quantile(0.95, sum by (le) (rate(trickster_proxy_request_duration_seconds_bucket{backend_name=\"greptimedb1\",provider=\"greptimedb\",path=\"query\"}[$__rate_interval])))", + "legendFormat": "p95", + "range": true, + "refId": "A" + } + ], + "title": "GreptimeDB Request Duration (p95)", + "type": "timeseries" + }, + { + "datasource": { + "type": "prometheus", + "uid": "ds_prom_direct" + }, + "fieldConfig": { + "defaults": { + "unit": "ops" + }, + "overrides": [] + }, + "gridPos": { + "h": 9, + "w": 12, + "x": 0, + "y": 91 + }, + "id": 11, + "options": { + "legend": { + "calcs": [], + "displayMode": "list", + "placement": "bottom", + "showLegend": true + }, + "tooltip": { + "mode": "multi", + "sort": "none" + } + }, + "targets": [ + { + "datasource": { + "type": "prometheus", + "uid": "ds_prom_direct" + }, + "editorMode": "code", + "expr": "sum by (cache_status) (rate(trickster_proxy_points_total{backend_name=\"greptimedb1\",provider=\"greptimedb\",path=\"query\"}[$__rate_interval]))", + "legendFormat": "{{cache_status}}", + "range": true, + "refId": "A" + } + ], + "title": "Returned Elements", + "type": "timeseries" + }, + { + "datasource": { + "type": "prometheus", + "uid": "ds_prom_direct" + }, + "fieldConfig": { + "defaults": { + "unit": "ops" + }, + "overrides": [] + }, + "gridPos": { + "h": 9, + "w": 12, + "x": 12, + "y": 91 + }, + "id": 12, + "options": { + "legend": { + "calcs": [], + "displayMode": "list", + "placement": "bottom", + "showLegend": true + }, + "tooltip": { + "mode": "multi", + "sort": "none" + } + }, + "targets": [ + { + "datasource": { + "type": "prometheus", + "uid": "ds_prom_direct" + }, + "editorMode": "code", + "expr": "sum by (operation, status) (rate(trickster_cache_operation_objects_total{cache_name=\"greptimedb_fs\"}[$__rate_interval]))", + "legendFormat": "{{operation}} / {{status}}", + "range": true, + "refId": "A" + } + ], + "title": "Cache Operations", + "type": "timeseries" + }, + { + "datasource": { + "type": "prometheus", + "uid": "ds_prom_direct" + }, + "fieldConfig": { + "defaults": { + "unit": "percent" + }, + "overrides": [] + }, + "gridPos": { + "h": 9, + "w": 12, + "x": 0, + "y": 100 + }, + "id": 13, + "options": { + "legend": { + "calcs": [], + "displayMode": "list", + "placement": "bottom", + "showLegend": true + }, + "tooltip": { + "mode": "multi", + "sort": "none" + } + }, + "targets": [ + { + "datasource": { + "type": "prometheus", + "uid": "ds_prom_direct" + }, + "editorMode": "code", + "expr": "100 * trickster_cache_usage_bytes{cache_name=\"greptimedb_fs\"} / trickster_cache_max_usage_bytes{cache_name=\"greptimedb_fs\"}", + "legendFormat": "{{cache_name}}", + "range": true, + "refId": "A" + } + ], + "title": "Cache Storage Utilization", + "type": "timeseries" + }, + { + "datasource": { + "type": "prometheus", + "uid": "ds_prom_direct" + }, + "fieldConfig": { + "defaults": { + "unit": "ops" + }, + "overrides": [] + }, + "gridPos": { + "h": 9, + "w": 12, + "x": 12, + "y": 100 + }, + "id": 14, + "options": { + "legend": { + "calcs": [], + "displayMode": "list", + "placement": "bottom", + "showLegend": true + }, + "tooltip": { + "mode": "multi", + "sort": "none" + } + }, + "targets": [ + { + "datasource": { + "type": "prometheus", + "uid": "ds_prom_direct" + }, + "editorMode": "code", + "expr": "sum by (reason) (rate(trickster_cache_events_total{cache_name=\"greptimedb_fs\",event=\"eviction\"}[$__rate_interval]))", + "legendFormat": "{{reason}}", + "range": true, + "refId": "A" + } + ], + "title": "Cache Evictions", + "type": "timeseries" + } + ], + "preload": false, + "refresh": "5s", + "schemaVersion": 41, + "tags": [ + "trickster", + "timescaledb" + ], + "templating": { + "list": [ + { + "allowCustomValue": false, + "current": { + "text": "greptimedb-direct", + "value": "ds_greptimedb_direct" + }, + "description": "", + "name": "datasource", + "options": [], + "query": "grafana-postgresql-datasource", + "refresh": 1, + "regex": "/^greptimedb-/", + "type": "datasource" + } + ] + }, + "time": { + "from": "now-2d", + "to": "now" + }, + "timepicker": {}, + "timezone": "utc", + "title": "GreptimeDB", + "uid": "trickster-greptimedb", + "version": 2 +} diff --git a/docs/developer/environment/docker-compose-data/grafana-config/provisioning/datasources/datasource.yaml b/docs/developer/environment/docker-compose-data/grafana-config/provisioning/datasources/datasource.yaml index 136a9bc9e..dd35c06aa 100644 --- a/docs/developer/environment/docker-compose-data/grafana-config/provisioning/datasources/datasource.yaml +++ b/docs/developer/environment/docker-compose-data/grafana-config/provisioning/datasources/datasource.yaml @@ -473,3 +473,72 @@ datasources: connMaxLifetime: 14400 secureJsonData: password: trickster-dev-grafana + +# GreptimeDB direct SQL, using Grafana's bundled PostgreSQL plugin. +- name: greptimedb-direct + type: grafana-postgresql-datasource + access: proxy + orgId: 1 + uid: ds_greptimedb_direct + url: greptimedb:4003 + user: grafana_ro + editable: true + isDefault: false + jsonData: + database: public + sslmode: disable + timescaledb: false + maxOpenConns: 5 + maxIdleConns: 2 + connMaxLifetime: 14400 + secureJsonData: + password: trickster-dev-grafana + +- name: greptimedb-prom-direct + type: prometheus + access: proxy + orgId: 1 + uid: ds_greptimedb_prom_direct + url: http://greptimedb:4000/v1/prometheus + basicAuth: true + basicAuthUser: grafana_ro + editable: true + isDefault: false + jsonData: + httpMethod: POST + secureJsonData: + basicAuthPassword: trickster-dev-grafana + +- name: greptimedb-trickster + type: grafana-postgresql-datasource + access: proxy + orgId: 1 + uid: ds_greptimedb_trickster + url: host.docker.internal:8489 + user: grafana_ro + editable: true + isDefault: false + jsonData: + database: public + sslmode: disable + timescaledb: false + maxOpenConns: 5 + maxIdleConns: 2 + connMaxLifetime: 14400 + secureJsonData: + password: trickster-dev-grafana + +- name: greptimedb-prom-trickster + type: prometheus + access: proxy + orgId: 1 + uid: ds_greptimedb_prom_trickster + url: http://host.docker.internal:8480/greptimedb1/v1/prometheus + basicAuth: true + basicAuthUser: grafana_ro + editable: true + isDefault: false + jsonData: + httpMethod: POST + secureJsonData: + basicAuthPassword: trickster-dev-grafana diff --git a/docs/developer/environment/docker-compose-data/greptimedb-config/standalone.toml b/docs/developer/environment/docker-compose-data/greptimedb-config/standalone.toml new file mode 100644 index 000000000..cf36d7e5d --- /dev/null +++ b/docs/developer/environment/docker-compose-data/greptimedb-config/standalone.toml @@ -0,0 +1,22 @@ +default_timezone = "UTC" +enable_telemetry = false + +[http] +addr = "0.0.0.0:4000" + +[grpc] +bind_addr = "0.0.0.0:4001" + +[mysql] +enable = true +addr = "0.0.0.0:4002" + +[postgres] +enable = true +addr = "0.0.0.0:4003" + +[influxdb] +enable = true + +[logging] +dir = "" diff --git a/docs/developer/environment/docker-compose-data/greptimedb-config/users b/docs/developer/environment/docker-compose-data/greptimedb-config/users new file mode 100644 index 000000000..4e07e3ad9 --- /dev/null +++ b/docs/developer/environment/docker-compose-data/greptimedb-config/users @@ -0,0 +1,4 @@ +# Public developer-environment credentials, not production accounts. +grafana_ro:ro=trickster-dev-grafana +trickster:ro=trickster-dev-upstream +seeder=trickster-dev-seed diff --git a/docs/developer/environment/docker-compose-data/prometheus-config/prometheus.yml b/docs/developer/environment/docker-compose-data/prometheus-config/prometheus.yml index beee29455..f104c86e1 100644 --- a/docs/developer/environment/docker-compose-data/prometheus-config/prometheus.yml +++ b/docs/developer/environment/docker-compose-data/prometheus-config/prometheus.yml @@ -23,6 +23,12 @@ rule_files: # - "first_rules.yml" # - "second_rules.yml" +remote_write: + - url: http://greptimedb:4000/v1/prometheus/write + basic_auth: + username: seeder + password: trickster-dev-seed + # A scrape configuration containing exactly one endpoint to scrape: # Here it's Prometheus itself. scrape_configs: diff --git a/docs/developer/environment/docker-compose.yml b/docs/developer/environment/docker-compose.yml index aadee1273..2a8e39af0 100644 --- a/docs/developer/environment/docker-compose.yml +++ b/docs/developer/environment/docker-compose.yml @@ -43,6 +43,8 @@ services: condition: service_healthy timescaledb: condition: service_healthy + greptimedb: + condition: service_healthy druid: condition: service_healthy @@ -206,7 +208,7 @@ services: restart: always # generates the shared synthetic trips seed data (hack/seedgen) that the - # clickhouse, mysql, timescaledb and druid seeders load. The output is + # clickhouse, mysql, timescaledb, greptimedb and druid seeders load. The output is # byte-identical on every run, is verified against a pinned hash, and needs # no network access, so the service runs with networking disabled to prove it. seed_data_generate: @@ -378,6 +380,59 @@ services: seed_data_generate: condition: service_completed_successfully + # GreptimeDB standalone, with HTTP, gRPC, MySQL and PostgreSQL endpoints. + greptimedb: + # Includes #9295 (Grafana's comment-only PostgreSQL health query). + image: greptime/greptimedb-nightly:nightly-20260923-e91faa9df@sha256:657e0d3e5899f7e17db2bf6810264c1d1e8fff84055a90938201aedc86a19468 + extra_hosts: + - "host.docker.internal:host-gateway" + command: + - standalone + - start + - --config-file=/etc/greptimedb/standalone.toml + - --data-home=/greptimedb-data + - --user-provider=static_user_provider:file:/etc/greptimedb/users + volumes: + - ./docker-compose-data/greptimedb-config:/etc/greptimedb:ro + - greptimedb-data:/greptimedb-data + ports: + - 4000:4000 + - 4001:4001 + - 4002:4002 + - 4003:4003 + restart: always + healthcheck: + test: ["CMD", "curl", "--fail", "--silent", "http://127.0.0.1:4000/health"] + interval: 5s + timeout: 3s + retries: 30 + start_period: 15s + + # SQL loader preserves the shared fixture's DATE and narrow numeric types. + greptimedb_seed: + image: golang:1.27-alpine + working_dir: /src + environment: + GREPTIMEDB_URL: http://greptimedb:4000 + GREPTIMEDB_DATABASE: public + GREPTIMEDB_SEED_USER: seeder + GREPTIMEDB_SEED_PASSWORD: trickster-dev-seed + GREPTIMEDB_SEED_DATA: /seed-data + GOFLAGS: -mod=mod + GOPROXY: "off" + GOTOOLCHAIN: local + GOCACHE: /go-cache + volumes: + - ../../../hack/greptimeseed:/src:ro + - ./docker-compose-data/seed-data:/seed-data:ro + - seedgen-gocache:/go-cache + entrypoint: ["go", "run", "."] + depends_on: + greptimedb: + condition: service_healthy + seed_data_generate: + condition: service_completed_successfully + # graphite (carbon-cache + graphite-web, whisper storage) graphite: image: graphiteapp/graphite-statsd:1.1.10-5 @@ -545,4 +600,5 @@ volumes: prometheus-seed-om: mysql-data: timescaledb-data: + greptimedb-data: redis-data: diff --git a/docs/developer/environment/trickster-config/trickster.yaml b/docs/developer/environment/trickster-config/trickster.yaml index e78e47bc4..f1af10052 100644 --- a/docs/developer/environment/trickster-config/trickster.yaml +++ b/docs/developer/environment/trickster-config/trickster.yaml @@ -9,6 +9,12 @@ listeners: timescaledb1: protocol: postgres port: 8488 + greptimedb1: + protocol: postgres + port: 8489 + greptimedb-mysql: + protocol: mysql + port: 8491 # clickhouse-native exposes click1 through the ClickHouse binary protocol. clickhouse-native: protocol: clickhouse @@ -28,6 +34,12 @@ authenticators: provider: basic users: grafana_ro: trickster-dev-grafana + greptimedb-grafana: + provider: basic + # HTTP still authenticates with the origin; pgwire uses upstream_url. + proxy_preserve: true + users: + grafana_ro: trickster-dev-grafana negative_caches: default: '400': 3s @@ -43,6 +55,8 @@ caches: provider: filesystem graphite_fs: provider: filesystem + greptimedb_fs: + provider: filesystem redis: provider: redis redis: @@ -132,6 +146,21 @@ backends: healthcheck: interval: 5s timeout: 3s + # One backend serves HTTP, pgwire and MySQL with separate native URLs. + greptimedb1: + provider: greptimedb + origin_url: 'http://127.0.0.1:4000' + listener_names: [default, greptimedb1, greptimedb-mysql] + authenticator_name: greptimedb-grafana + # A dedicated filesystem cache exposes storage gauges for this dashboard. + cache_name: greptimedb_fs + postgres: + upstream_url: 'postgres://trickster:trickster-dev-upstream@127.0.0.1:4003/public' + mysql: + upstream_url: 'mysql://trickster:trickster-dev-upstream@127.0.0.1:4002/public' + healthcheck: + interval: 5s + timeout: 3s click1: provider: clickhouse origin_url: 'http://127.0.0.1:8123' diff --git a/docs/developer/sql-dialect-adapters.md b/docs/developer/sql-dialect-adapters.md index 5b5d7832c..e0b70939f 100644 --- a/docs/developer/sql-dialect-adapters.md +++ b/docs/developer/sql-dialect-adapters.md @@ -1,5 +1,10 @@ # Adding a SQL Dialect Adapter +For another database using PostgreSQL's wire protocol, also follow +[Adding a PostgreSQL Wire Engine](adding-pgwire-engine.md). It covers shared +transport, session and time-axis hooks, multi-protocol providers and the +compatibility-corpus runner. + Trickster accelerates SQL-based time series backends (currently ClickHouse and MySQL) by parsing each query into a dialect-native abstract syntax tree, analyzing it diff --git a/docs/developer/streaming-unmarshalers.md b/docs/developer/streaming-unmarshalers.md new file mode 100644 index 000000000..06c443584 --- /dev/null +++ b/docs/developer/streaming-unmarshalers.md @@ -0,0 +1,218 @@ +# Streaming Unmarshalers + +A time series backend's [Modeler](../../pkg/timeseries/modeler.go) turns each upstream response body into Trickster's Common Time Series Format, the [`dataset.DataSet`](../../pkg/timeseries/dataset/dataset.go). Most providers do this in two steps. They unmarshal the whole body into a provider-specific model, such as a set of JSON structs, and then copy that model into a DataSet. That holds two full copies of the data in memory at once and allocates heavily along the way. + +The packages described here let a provider decode a response in one pass, straight into a DataSet, with no intermediate model. Rows and points may arrive in any order, so a provider never needs to rewrite upstream queries (for example, by adding `ORDER BY`) to make decoding work. + +| Package | Provides | +| --- | --- | +| [`pkg/timeseries/dataset`](../../pkg/timeseries/dataset/builder.go) | `Builder`, which assembles a DataSet from rows or points in any order | +| [`pkg/timeseries/dataset/stream`](../../pkg/timeseries/dataset/stream/stream.go) | the `Decoder` interface, line and JSON decoders, value parsers, and Modeler adapters | +| [`pkg/timeseries/dataset/stream/streamtest`](../../pkg/timeseries/dataset/stream/streamtest/streamtest.go) | conformance checks and benchmarks for decoders | +| [`pkg/timeseries/epoch`](../../pkg/timeseries/epoch/parse.go) | `ParseDecimal`, exact parsing of numeric timestamps | + +## How It Fits Together + +A provider writes a `stream.NewDecoderFunc`. Given the request's `TimeRangeQuery`, it returns a `stream.Decoder` for one response. A `Decoder` is an `io.Writer` and an `io.ReaderFrom`, plus a `Finish` method that returns the decoded `timeseries.Timeseries`. The stream package's adapters turn that function into the Modeler's wire unmarshalers: + +```go +func NewModeler() *timeseries.Modeler { + return timeseries.NewModeler( + stream.BytesUnmarshaler(newDecoder), stream.ReaderUnmarshaler(newDecoder), + MarshalTimeseries, MarshalTimeseriesWriter, + dataset.UnmarshalDataSet, dataset.MarshalDataSet) +} +``` + +`ReaderUnmarshaler` passes the response body to the decoder's `ReadFrom`, so decoding happens while the body is read. If you feed a decoder yourself, call `ReadFrom` instead of using `io.Copy`. When the source is a `bytes.Reader`, `io.Copy` uses the reader's `WriteTo`, which delivers the whole body in a single `Write`, and the JSON decoder then has to buffer all of it before it can start. + +Today the proxy engine reads each upstream body into memory before calling the unmarshaler, so the current savings come from skipping the intermediate model. Once the engine passes response bodies through directly, the same decoders will read from the network with no changes. + +## Choosing a Decoder + +### Newline-Delimited Formats + +`stream.NewLines(onLine, finish)` handles formats with one record per line, such as TSV, CSV without quoted newlines, and JSON Lines. It calls `onLine` for each line without its `\n` or `\r\n` terminator, even when the line was split across `Write` calls, and delivers a final unterminated line during `Finish`. The line is only valid during the call. Lines longer than 16 MiB fail with `ErrLineTooLong` before they are buffered, so an overlong line cannot grow memory; `SetMaxLineBytes` changes the limit. + +`stream.SplitFields(line, sep, dst)` splits a line into fields without allocating, reusing `dst`. It does not interpret quotes or escapes, so unescape fields yourself where the format requires it. + +### JSON Documents + +`stream.NewJSON(walk, finish)` handles a single JSON document. The `walk` function receives an `encoding/json` `Decoder` with `UseNumber` set, and must consume exactly one value from it. Only whitespace may follow that value: a second JSON value fails with `ErrTrailingData`, and anything else fails as a syntax error. Three helpers let a walk hold only the current token or element in memory: + +- `stream.Object(dec, func(key string) error)` calls the function for each key, in the order the keys arrive. +- `stream.Array(dec, func() error)` calls the function once for each element. +- `stream.Skip(dec)` consumes and discards the next value a token at a time, so a large skipped value is never held in memory. + +Each callback must consume the value it was called for, for example with `dec.Decode`, a nested `Object` or `Array`, or `Skip`. A callback that returns without doing so fails with `ErrValueNotConsumed`. `Object` and `Array` return `ErrNull` for a JSON `null` after consuming it, so a caller that accepts a null can check with `errors.Is` and carry on. They return `ErrUnexpectedToken` when the value is the wrong kind. + +Decode every row into the same `[]json.RawMessage`. `encoding/json` reuses the slice and each element's buffer, so decoding rows stops allocating once those buffers have grown. + +JSON does not guarantee key order. If something you need first, such as a schema, might arrive after the data that depends on it, hold the early data as a `json.RawMessage` and process it once the schema has been read. + +### Other Formats + +A format that fits neither decoder can implement `stream.Decoder` directly. The conformance checks feed a decoder with `Write` calls, with a single `ReadFrom` call, or with `Write` calls followed by one `ReadFrom`, and expect the same result each way. Errors should be sticky, and `Finish` is called once. + +## Building the DataSet + +`dataset.NewBuilder(trq, opts)` returns a Builder for one response. `BuilderOptions` sets: + +- `Fields`: the timestamp, tag and value fields of each row. +- `SeriesName` and `QueryStatement`: copied into each series header the Builder creates. +- `Duplicates`: what to do with points in one series that share an epoch: `DuplicatesKeep`, `DuplicatesFirstWins`, `DuplicatesLastWins` or `DuplicatesError`. +- `SortSeries`: sorts each result's series by their tags when the build finishes. +- `TagString`: converts a tag's raw bytes to its value in the series' `Tags`. By default the bytes are used as they are; `stream.JSONTagString` unquotes JSON strings. + +### Row Mode + +Use row mode for formats that send one row per point, such as SQL results and TSV or CSV: + +```go +r := b.Row() // reused, and valid until the next call to Row +r.SetEpoch(ep) +r.SetTag(0, host) // an index into BuilderOptions.Fields.Tags; the bytes are copied +r.AddValue(v) // in BuilderOptions.Fields.Values order +if err := r.Commit(); err != nil { + return err +} +``` + +The Builder remembers each raw tag encoding it has seen, so a row that repeats an earlier row's tag bytes finds its series with one lookup that does not allocate. A new encoding is converted with `TagString` and matched against the existing series by header, so equivalent encodings, such as `"a"` and `"\u0061"` in JSON, share a series. A tag that is never set is left out of the series' `Tags`, so an unset tag and an empty one produce different series. + +### Series Mode + +Use series mode for formats that send each series as one block, such as a Prometheus matrix or InfluxQL JSON: + +```go +b.StartSeries(dataset.SeriesHeader{Name: "up", Tags: tags, ValueFieldsList: valueFields}) +for _, p := range points { + r := b.Row() + r.SetEpoch(p.epoch) + r.AddValue(p.value) + if err := r.Commit(); err != nil { + return err + } +} +b.EndSeries() +``` + +Rows committed while a series is open go to that series and may not set tags. `StartSeries` reopens the series with an identical header if there is one, so a series that arrives in pieces becomes one series. `AppendPoint` adds a `Point` you have already built. For formats that return several statements, `SetResult(statementID, name)` sends later rows and series to another result, creating it if needed. + +### Finishing + +The Builder matches series the same way merges do: the header hash finds candidates, and a comparison of the headers confirms the match. So the DataSet never holds two series that a later merge would treat as one, and two different series whose hashes collide stay separate, both in the Builder and in later merges. + +`Finish` returns the DataSet. It sorts only the series whose points arrived out of order, using a stable sort that keeps arrival order among equal epochs, and then applies the duplicate policy. When a series' points do arrive in order, duplicates are handled as they arrive, so `DuplicatesError` fails the `Commit` immediately. `Finish` also calculates each series header's size, and sets the DataSet's `TimeRangeQuery` and `ExtentList` from the query. + +Point values are carved from shared, chunked backing arrays, so a point does not need an allocation of its own. `dataset.PointSize` is the size estimate the Builder records for each point. + +`ErrInvalidRow` and `ErrDuplicateEpoch` wrap `timeseries.ErrInvalidBody`, and `ErrBuilderFinished` reports use after `Finish`. `ErrInvalidRow` covers: + +- a row without an epoch; +- a row with the wrong number of values; +- a tag index out of range; +- tags set on a series-mode row; +- `AppendPoint` with no open series. + +`ErrDuplicateEpoch` reports a duplicate under `DuplicatesError`. + +A provider that wraps the DataSet in its own `Timeseries` type can do so in its finish function, because a `FinishFunc` returns a `timeseries.Timeseries`. + +## Parsing Values + +- `epoch.ParseDecimal(raw, unit)` parses a count of seconds, milliseconds, microseconds or nanoseconds, where `unit` is `DateTimeUnixSecs`, `DateTimeUnixMilli`, `DateTimeUnixMicro` or `DateTimeUnixNano`. The count may have a fraction and an exponent, and may be wrapped in JSON quotes. It uses integer math, and rejects precision finer than one nanosecond and values outside the `int64` range. +- `stream.ParseValue(raw, dt)` parses the text of a TSV or CSV cell as the field's data type: + - integer and Unix timestamp types return `int64`, except `Uint64`, which returns `uint64`; + - `Float64` returns `float64`, and `Bool` returns `bool`; + - text types, and RFC 3339 and SQL date and time types, return `string`; + - empty text is `nil` for any type except text; + - `Unknown` infers a `bool` or number from JSON-style literals and falls back to `string`. +- `stream.ParseJSONValue(raw, dt)` parses a raw JSON value. `null` is `nil`, and quoted values are unquoted first, so numbers that an API sends as strings still parse as numbers. + +Parse errors wrap `stream.ErrInvalidValue`, which wraps `timeseries.ErrInvalidBody`. + +## Testing a Decoder + +`streamtest.Conformance(t, newDecoder, c)` feeds `c.Body` to the decoder every way it can be fed: +- in one `Write`, one byte at a time, and in random chunks; +- through `ReadFrom` with several reader behaviors; +- as `Write` calls followed by `ReadFrom`; +- through both adapters. + +It reports each result that differs, and checks that a failed read surfaces as an error instead of a partial result. The `Case` adds optional checks: + +- `WantErr`: the error every feed must return, matched with `errors.Is`. `streamtest.ErrAny` accepts any error. +- `Want`: the DataSet every feed must produce, ignoring sizes. +- `Legacy`: an existing unmarshaler whose DataSet must match, ignoring sizes. +- `Shuffle`: reorders the body without changing its meaning; the result must match apart from series order. `streamtest.ShuffleLines(n)` shuffles every line after the first `n`. +- `Unwrap`: extracts the DataSet from a provider's wrapper type. + +`streamtest.Compare(want, got, opts)` reports the first difference between two DataSets, for your own assertions. `streamtest.Bench` measures an unmarshaler the way the proxy engine calls it, so an old and a new decoder can be compared side by side: + +```go +b.Run("legacy", func(b *testing.B) { streamtest.Bench(b, model.UnmarshalTimeseriesReader, trq, body) }) +b.Run("stream", func(b *testing.B) { streamtest.Bench(b, stream.ReaderUnmarshaler(newDecoder), trq, body) }) +``` + +Randomness in these tests comes from `pkg/util/weak/weaktest`; see [Non-Cryptographic Randomness](./weak-randomness.md). + +## A Complete Decoder + +This decoder reads `time`, `host` and `value` columns from tab-separated rows that may arrive in any order: + +```go +var fields = timeseries.SeriesFields{ + Timestamp: timeseries.FieldDefinition{Name: "time", DataType: timeseries.DateTimeUnixMilli, + Role: timeseries.RoleTimestamp}, + Tags: timeseries.FieldDefinitions{{Name: "host", DataType: timeseries.String, + Role: timeseries.RoleTag, OutputPosition: 1}}, + Values: timeseries.FieldDefinitions{{Name: "value", DataType: timeseries.Float64, + Role: timeseries.RoleValue, OutputPosition: 2}}, +} + +func newTSVDecoder(trq *timeseries.TimeRangeQuery) (stream.Decoder, error) { + b := dataset.NewBuilder(trq, dataset.BuilderOptions{ + Fields: fields, SeriesName: "tsv", Duplicates: dataset.DuplicatesError, + }) + var header bool + var cols [][]byte + onLine := func(line []byte) error { + if !header { + header = true // the first line names the columns + return nil + } + cols = stream.SplitFields(line, '\t', cols) + if len(cols) != 3 { + return timeseries.ErrInvalidBody + } + ep, err := epoch.ParseDecimal(cols[0], timeseries.DateTimeUnixMilli) + if err != nil { + return err + } + v, err := stream.ParseValue(cols[2], timeseries.Float64) + if err != nil { + return err + } + r := b.Row() + r.SetEpoch(ep) + r.SetTag(0, cols[1]) + r.AddValue(v) + return r.Commit() + } + return stream.NewLines(onLine, func() (timeseries.Timeseries, error) { + return b.Finish() + }), nil +} +``` + +[`decoders_test.go`](../../pkg/timeseries/dataset/stream/decoders_test.go) in the stream package has this decoder with header validation, along with two more to start from. `newMatrixDecoder` decodes a Prometheus-style matrix in series mode, and `newRowsDecoder` decodes JSON rows and checks a row count that arrives after them. Their tests are in [`conformance_test.go`](../../pkg/timeseries/dataset/stream/conformance_test.go). + +## Converting an Existing Provider + +1. Write the decoder beside the provider's current unmarshaler. +2. Run `streamtest.Conformance` over the provider's test bodies, with `Legacy` set to the current `WireUnmarshalerReader`. Sizes may differ, but everything else should match. Where the old behavior is a bug, assert the corrected result with `Want` instead, and point it out in the pull request. +3. Compare the old and new decoders with `streamtest.Bench`. +4. Point the Modeler's wire unmarshalers at the adapters, and remove the old model code. + +Formats that send each series as one block, with its points in time order, map directly onto series mode; examples are Prometheus, InfluxQL JSON and Graphite. Formats that send rows in no particular order use row mode and rely on the Builder to sort when needed; examples are InfluxDB 3 SQL, ClickHouse, Flux CSV and Druid. The MySQL provider's wire-protocol path never builds a DataSet, so it is not a candidate. diff --git a/docs/developer/weak-randomness.md b/docs/developer/weak-randomness.md new file mode 100644 index 000000000..0d890aaac --- /dev/null +++ b/docs/developer/weak-randomness.md @@ -0,0 +1,76 @@ +# Non-Cryptographic Randomness + +Trickster uses fast, non-cryptographic ("weak") random numbers in a few places. The application uses them for request correlation IDs, traffic sampling and simulated latency. Tests use them for reproducible data and shuffles. A weak source is fine for these jobs, but it is predictable, so it must never produce anything an attacker should not be able to guess. + +To keep every use visible and reviewable, weak randomness goes through `pkg/util/weak`, with the narrow standalone-core exception below. The application draws from one package, tests from another, and neither may use the other's. This lets us audit production use, phase it out when a use becomes security-relevant, and keep test-only helpers out of the running application. + +## Which Package to Use + +| Where the code lives | Use | +| --- | --- | +| Application code: non-test files under `cmd/` and `pkg/` | [`pkg/util/weak/compat`](../../pkg/util/weak/compat/compat.go) | +| Tests and tooling: `_test.go` files, `pkg/testutil`, packages whose directory name ends in `test` (such as `streamtest`), and everything under `integration/`, `examples/` and `hack/` | [`pkg/util/weak/weaktest`](../../pkg/util/weak/weaktest/weaktest.go) | +| Anything security-relevant: tokens, keys, nonces, session IDs | `crypto/rand`, never `pkg/util/weak` | + +Do not introduce direct `math/rand` or `math/rand/v2` imports elsewhere. The standalone `pkg/lb` tree has a separately tested standard-library-only dependency boundary. Its four existing `math/rand/v2` users are explicitly allowed: `hrw/hrw.go` and `p2c/p2c.go` for load spreading, `rr/rr.go` for the initial rotation offset, and `lbtest/lbtest.go` for reproducible test flows. Their existing `gosec` justifications remain in place; the exception does not cover new files or cryptographic uses. + +### Application Code + +`compat` offers `Uint64`, `IntN` and `Int64`, all drawn from a randomly seeded source: + +```go +import "github.com/trickstercache/trickster/v2/pkg/util/weak/compat" + +if compat.IntN(100) < percent { + // mirror this request +} +``` + +Keep `compat` small. When adding a caller, say in the pull request why a predictable value is acceptable there. When a use moves to `crypto/rand`, remove any `compat` function that no longer has callers. + +### Tests + +`weaktest.NewRand` returns a generator whose sequence is fixed by its two seeds, so a failing test fails the same way every run. `weaktest.IntN` draws from a randomly seeded source, for tests that do not care which values they get: + +```go +import "github.com/trickstercache/trickster/v2/pkg/util/weak/weaktest" + +rng := weaktest.NewRand(42, 0) // same seeds, same sequence +rng.Shuffle(len(rows), func(i, j int) { rows[i], rows[j] = rows[j], rows[i] }) + +n := weaktest.IntN(1000) +``` + +`weaktest.Rand` is an alias for the `math/rand/v2` generator. Helpers can accept a `*weaktest.Rand` and call its methods without importing `math/rand/v2`. + +## How the Rules Are Enforced + +The rules are checked at runtime and at build time, so a defect in one check is caught by another. + +### At Runtime + +The application's `main` calls `weak.RegisterUnlessTestMode()` before anything else. From then on, every `weaktest` function panics with `weak: test-only randomness used outside of a test; use package compat`. Unit tests never run `main`, so `weaktest` works normally inside them. + +`compat` has no runtime check. Application code such as mirror sampling runs inside unit tests too, and a process-wide switch cannot tell that code apart from test code that misuses `compat`. The build-time checks below cover that direction. + +### `make check-weak-random` + +[`hack/check-weak-random`](../../hack/check-weak-random/main.go) reads the imports of every Go file under `cmd/`, `pkg/`, `integration/`, `examples/` and `hack/`, and fails when: + +- a file outside `pkg/util/weak/` imports `math/rand` or `math/rand/v2`, except for the four standalone-core uses listed above +- test or tooling code imports `pkg/util/weak/compat` +- application code imports `pkg/util/weak/weaktest` + +Both `make test`, which CI runs, and `make lint` run this check. Its own unit tests also check the whole repository, so a plain `go test` catches violations as well. + +### `golangci-lint` + +The `forbidigo` linter in [`.golangci.yml`](../../.golangci.yml) flags any use of a `math/rand` or `math/rand/v2` function or type outside `pkg/util/weak/` and the four standalone-core files. It resolves imports by type, so renamed imports such as `mrand "math/rand/v2"` are caught too. Methods on a `*weaktest.Rand` are still allowed. The configuration does not lint `_test.go` files, so `make check-weak-random` is what covers tests. `gosec` also still reports weak randomness (G404); its reviewed exceptions are confined to these locations. + +## Test Mode for Launched Processes + +A test that starts a real trickster process can set the `TRICKSTER_TEST_MODE` environment variable to any non-empty value. `main` then skips registration, and `weaktest` stays usable in that process. Never set it in production. Integration tests that run trickster in-process through `daemon.Start` never call `main`, so they do not need it. + +## Adding a Test Helper Package + +A new helper package that tests import is test code, so it must use `weaktest`. Name its directory with a `test` suffix, like `streamtest`, or place it under `pkg/testutil`, so `make check-weak-random` classifies it as test code. diff --git a/docs/greptimedb.md b/docs/greptimedb.md new file mode 100644 index 000000000..9d6b41fbc --- /dev/null +++ b/docs/greptimedb.md @@ -0,0 +1,182 @@ +# GreptimeDB Provider + +Use `provider: greptimedb` for GreptimeDB's HTTP, PostgreSQL and MySQL query +surfaces. One backend may be mapped to all three listener protocols. The +provider reuses the SQL and Prometheus cache engines, with GreptimeDB-specific +parsing, session settings, result types and API paths. + +## Surfaces + +| Surface | Origin port | Behavior | +| --- | --- | --- | +| HTTP `/v1/sql` | 4000 | Eligible SELECT results use delta or object caching | +| PostgreSQL wire | 4003 | Eligible simple queries use delta or object caching; extended queries are relayed | +| MySQL wire | 4002 | Eligible text queries use delta or object caching | +| `/v1/prometheus/api/v1/query_range` | 4000 | Prometheus range caching and supported time-series merges | +| Prometheus instant and metadata APIs | 4000 | Object caching subject to HTTP cache policy | +| Ingest, TQL, `/v1/promql`, other HTTP APIs | 4000 | Proxied without SQL/PromQL delta rewriting | +| gRPC | 4001 | Not implemented by this provider | + +Wire compatibility does not make GreptimeDB a PostgreSQL or MySQL server. +Use the `greptimedb` provider even on native listeners. For example, MySQL's +`SHOW COUNT(*) WARNINGS` is not a supported GreptimeDB statement, and its +PostgreSQL session settings have different semantics. + +The developer environment pins an official nightly that includes +[GreptimeDB #9295](https://github.com/GreptimeTeam/greptimedb/pull/9295), the fix +for Grafana's comment-only PostgreSQL health query. See the +[developer environment](developer/environment/README.md#greptimedb-details) +for the exact image and reproducible checks. The tested Grafana plugin is +the bundled PostgreSQL datasource, not a GreptimeDB-specific plugin. + +## Configuration + +The [complete example](../examples/conf/greptimedb.yaml) exposes HTTP on 8480, +PostgreSQL on 8489 and MySQL on 8491. A minimal mixed backend is: + +```yaml +listeners: + greptime-pg: + protocol: postgres + port: 8489 + greptime-mysql: + protocol: mysql + port: 8491 + +authenticators: + greptime-readers: + provider: basic + proxy_preserve: true + users: + grafana_ro: ${GRAFANA_RO_PASSWORD} + +backends: + greptime: + provider: greptimedb + origin_url: http://greptime.example:4000 + listener_names: [default, greptime-pg, greptime-mysql] + authenticator_name: greptime-readers + postgres: + upstream_url: postgres://trickster_ro:REPLACE_ME@greptime.example:4003/public + mysql: + upstream_url: mysql://trickster_ro:REPLACE_ME@greptime.example:4002/public +``` + +HTTP uses `origin_url`; `postgres.upstream_url` and `mysql.upstream_url` +separately identify native endpoints and credentials. Never put the HTTP +port in either native URL. Without a native override, an HTTP origin's host +is reused with the engine's native default port, not its HTTP port or userinfo. +Explicit URLs are recommended. URL credentials do not expand environment +variables; authenticator user values do. + +List only `default` for an HTTP-only backend, or just the chosen native +listener for a native-only backend. A mixed backend probes HTTP `/health`; +a native-only backend uses its native authenticated probe. Listener, origin, +credential or protocol changes restart and drain affected native listeners. + +### Authentication And TLS + +For HTTP, `proxy_preserve: true` forwards client credentials to GreptimeDB. +Authenticated object responses are not shared unless the origin explicitly +permits HTTP caching. SQL and PromQL cache identities retain their database, +request parameters, effective timezone and authentication partitioning. + +For PostgreSQL, an authenticator terminates client authentication at +Trickster; the separate upstream URL supplies the origin role. Without an +authenticator, authentication is passed through to GreptimeDB. MySQL requires +a configured authenticator. Give every origin role only the permissions its +clients need; terminating authentication does not preserve separate upstream +roles for each client. + +Native TLS and limits use the existing `postgres` and `mysql` option blocks; +see [PostgreSQL](postgres.md) and [MySQL](mysql.md). Support for an option in +Trickster is not a promise that a particular GreptimeDB deployment supports +the corresponding server feature. The loopback developer environment uses +plaintext and published test credentials and must not be exposed publicly. + +## SQL Cache Classification + +Delta caching requires one provably deterministic time-bucketed SELECT with +supported grouping, ordering and time bounds. Other supported reads can use +whole-result object caching. Writes, volatile queries and session-unsafe +statements bypass caches; a parser rejection never makes a write cacheable. + +PostgreSQL supports fixed-width `date_bin` intervals or compact widths such +as `5m`, UTC fixed-width `date_trunc`, and epoch buckets of the form +`floor(extract(epoch FROM ts)/300)*300` or `date_part`. Calendar widths, +ambiguous time predicates, unknown timezone-dependent semantics, and +submicrosecond PostgreSQL text buckets do not get delta plans. Extended +protocol messages are always relayed. + +HTTP SQL supports GET and form-encoded POST with the default `greptimedb_v1` +response format. A delta response retains typed schema, row ordering, NULLs +and exact integer values. Like PostgreSQL, unaligned SQL bounds round inward +to complete buckets before delta caching; partial edge buckets are omitted. +Alternate formats, `limit`, unknown options and ranges with no complete bucket +retain the original query. Rebuilt responses do not claim the origin's +execution duration or execution metrics. The default JSON response is decoded +row by row into a DataSet, and its validated client serialization is reused. + +MySQL supports `DATE_BIN('1m', ts, FROM_UNIXTIME(0))` and fixed-width +`DATE_TRUNC` buckets with verified UTC sessions, whole-second cadence and +half-open bounds, rounded inward to complete buckets. Other bucket origins, +subsecond cadence, inclusive upper bounds and ranges with no complete bucket +use the original query instead. Timestamp results keep up to nine +fractional digits, text groups are compared case-sensitively, and NULL ordering +follows GreptimeDB. Unsupported or failed session changes conservatively +disable caching for that connection. + +PromQL range endpoints round down to epoch-aligned steps, including 500ms +steps, so equivalent aligned ranges share cached samples. Instant and metadata +timestamps are unchanged. Route defaults do not grant permission to share +authenticated object responses: only an origin's explicit cache policy can +authorize that sharing. + +### Grafana Macros + +Use the built-in PostgreSQL datasource with **TimescaleDB disabled** and a +minimum interval of one minute. Native epoch-second columns are integers, +not timestamps cast to integers (GreptimeDB timestamp casts can yield +nanoseconds). + +| Macro family | GreptimeDB behavior | +| --- | --- | +| `$__time`, `$__timeEpoch` | Supported projection; not a bucket by itself | +| `$__timeFilter`, `$__timeFrom`, `$__timeTo` | Supported timestamp bounds | +| `$__timeGroup`, `$__timeGroupAlias` | Fixed epoch-floor buckets with TimescaleDB mode off | +| `$__unixEpochFilter`, `$__unixEpochFrom`, `$__unixEpochTo` | Numeric epoch-second bounds | +| `$__unixEpochGroup`, `$__unixEpochGroupAlias` | Fixed buckets over an epoch-second column | +| `$__unixEpochNanoFilter` | Supported SQL predicate; not an automatic delta plan | +| `$__interval` | Use a positive supported fixed duration | + +The [compatibility corpus](../pkg/backends/greptimedb/testdata/compatibility/README.md) +records real Grafana expansions, including unaligned windows and conservative +fallbacks. Grafana's MySQL macros are not interchangeable: `UNIX_TIMESTAMP` +is unsupported in the tested GreptimeDB image. Use explicit GreptimeDB SQL +when connecting a MySQL client. + +## Known Limits + +- `RANGE ... ALIGN` remains whole-result caching; TQL is passthrough. There + is no advanced RANGE/FILL dependency rewriting. +- PostgreSQL text timestamps carry microseconds even for a TIMESTAMP(9) + column. MySQL and HTTP have separate precision contracts. +- Transaction commands accepted by GreptimeDB are compatibility stubs, not + proof of transactional isolation. PostgreSQL cancellation is not a supported + upstream capability in the tested image. +- `SET LOCAL` persists at session scope in GreptimeDB. Effective settings are + probed and tracked; requested startup parameters alone are not proof of the + resulting timezone. The development read-only account cannot change it. +- PostgreSQL session load balancing is unsupported. MySQL listeners use the + existing native MySQL routing rules. +- PromQL `count_values` currently loses its grouping label upstream; merges + reject that response rather than claim a correct aggregate. +- Ingestion, mutations, enterprise-only features and gRPC caching are outside + this provider's accelerated contract. + +Run the [acceptance suites](../integration/greptimedb/README.md) against an +isolated seeded environment. Inspect cache counters as well as results: +`trickster_sql_query_analysis_total`, `trickster_sql_query_cache_total`, and +`trickster_sql_query_rewrite_failures_total` distinguish actual cache hits, +fallbacks and failed rewrites. SQL dialect labels are `greptimedb`; transport +metrics retain the native protocol name. diff --git a/docs/supported-backend-providers.md b/docs/supported-backend-providers.md index 893b1c6e9..aa951baf1 100644 --- a/docs/supported-backend-providers.md +++ b/docs/supported-backend-providers.md @@ -59,6 +59,14 @@ See the [PostgreSQL and TimescaleDB Provider Guide](./postgres.md) for the supported clients, SQL, authentication, TLS, caching, routing, and operations contract. +### GreptimeDB + +Trickster accelerates eligible GreptimeDB SQL queries over HTTP, PostgreSQL +and MySQL, plus its Prometheus-compatible range API. Specify `greptimedb` as +the provider and map native listeners explicitly. See the +[GreptimeDB Provider Guide](./greptimedb.md) for configuration, cache eligibility, +Grafana macros, authentication and upstream compatibility limits. + ### MySQL Trickster supports protocol-aware acceleration for supported MySQL diff --git a/examples/conf/example.full.yaml b/examples/conf/example.full.yaml index bdff3512f..9c3791e7b 100644 --- a/examples/conf/example.full.yaml +++ b/examples/conf/example.full.yaml @@ -490,6 +490,19 @@ backends: # # timeseries_retention_factor, backfill_tolerance, and # # max_object_size_bytes backend settings apply. + # # GreptimeDB serves HTTP and either or both native protocols. See + # # examples/conf/greptimedb.yaml for the corresponding listeners/authenticator. + # greptimedb1: + # provider: greptimedb + # origin_url: http://greptime.example:4000 + # listener_names: [default, greptime-pg, greptime-mysql] + # authenticator_name: greptime-readers + # cache_name: default + # postgres: + # upstream_url: postgres://readonly:REPLACE_ME@greptime.example:4003/public + # mysql: + # upstream_url: mysql://readonly:REPLACE_ME@greptime.example:4002/public + # # # example postgres backend, exposed by a listener with protocol postgres # # (see the listeners section). Provider timescaledb is an alias of postgres. # postgres1: diff --git a/examples/conf/greptimedb.yaml b/examples/conf/greptimedb.yaml new file mode 100644 index 000000000..772250a3d --- /dev/null +++ b/examples/conf/greptimedb.yaml @@ -0,0 +1,31 @@ +# Replace credentials and enable deployment-appropriate TLS before use. +listeners: + greptime-pg: + protocol: postgres + port: 8489 + greptime-mysql: + protocol: mysql + port: 8491 + +authenticators: + greptime-readers: + provider: basic + proxy_preserve: true + users: + grafana_ro: ${GRAFANA_RO_PASSWORD} + +backends: + greptime: + provider: greptimedb + origin_url: http://greptime.example:4000 + listener_names: [default, greptime-pg, greptime-mysql] + authenticator_name: greptime-readers + cache_name: default + postgres: + upstream_url: postgres://trickster_ro:REPLACE_ME@greptime.example:4003/public + upstream_tls_mode: disable + mysql: + upstream_url: mysql://trickster_ro:REPLACE_ME@greptime.example:4002/public + healthcheck: + interval: 5s + timeout: 3s diff --git a/hack/check-weak-random/main.go b/hack/check-weak-random/main.go new file mode 100644 index 000000000..46b68bc1f --- /dev/null +++ b/hack/check-weak-random/main.go @@ -0,0 +1,129 @@ +// Command check-weak-random fails when non-cryptographic randomness crosses the +// line between application code and test code; see package pkg/util/weak. +package main + +import ( + "errors" + "fmt" + "go/parser" + "go/token" + "io/fs" + "os" + "path" + "strconv" + "strings" +) + +const ( + weakDir = "pkg/util/weak/" + compatPath = "github.com/trickstercache/trickster/v2/pkg/util/weak/compat" + weaktestPath = "github.com/trickstercache/trickster/v2/pkg/util/weak/weaktest" +) + +var roots = []string{"cmd", "pkg", "integration", "examples", "hack"} + +func main() { + problems, err := check(".") + if err != nil { + if _, err2 := fmt.Fprintln(os.Stderr, err); err2 != nil { + fmt.Printf("failed to print error to STDERR: %s (%s)", err, err2) + } + os.Exit(2) + } + if len(problems) > 0 { + fmt.Print("Non-cryptographic randomness must go through pkg/util/weak/compat in the\n" + + "application and pkg/util/weak/weaktest in tests and tooling:\n\n") + for _, p := range problems { + fmt.Println(p) + } + fmt.Println() + os.Exit(1) + } + fmt.Print("\nWeak randomness uses the reviewed wrappers and standalone-core exceptions.\n\n") +} + +func check(repoRoot string) ([]string, error) { + r, err := os.OpenRoot(repoRoot) + if err != nil { + return nil, err + } + defer r.Close() + fsys := r.FS() + var problems []string + for _, root := range roots { + err := fs.WalkDir(fsys, root, func(p string, d fs.DirEntry, err error) error { + if err != nil { + return err + } + if d.IsDir() { + if d.Name() == "vendor" || d.Name() == "testdata" { + return fs.SkipDir + } + return nil + } + if path.Ext(p) != ".go" { + return nil + } + src, err := fs.ReadFile(fsys, p) + if err != nil { + return err + } + found, err := violations(p, src) + problems = append(problems, found...) + return err + }) + if err != nil && !errors.Is(err, fs.ErrNotExist) { + return nil, err + } + } + return problems, nil +} + +func violations(file string, src []byte) ([]string, error) { + // the weak packages and their tests are the one place that may use both sides + if strings.HasPrefix(file, weakDir) { + return nil, nil + } + f, err := parser.ParseFile(token.NewFileSet(), file, src, parser.ImportsOnly) + if err != nil { + return nil, err + } + app := isApplication(file) + var out []string + for _, spec := range f.Imports { + imp, err := strconv.Unquote(spec.Path.Value) + if err != nil { + return nil, err + } + switch { + case imp == "math/rand/v2" && standaloneRandom(file): + // These pre-existing uses keep the load-balancing core independently importable. + case imp == "math/rand" || imp == "math/rand/v2": + out = append(out, file+": imports "+imp+"; use pkg/util/weak/compat or pkg/util/weak/weaktest") + case imp == compatPath && !app: + out = append(out, file+": test or tooling code imports pkg/util/weak/compat; use pkg/util/weak/weaktest") + case imp == weaktestPath && app: + out = append(out, file+": application code imports pkg/util/weak/weaktest; use pkg/util/weak/compat") + } + } + return out, nil +} + +func standaloneRandom(file string) bool { + switch file { + case "pkg/lb/hrw/hrw.go", "pkg/lb/p2c/p2c.go", "pkg/lb/rr/rr.go", "pkg/lb/lbtest/lbtest.go": + return true + } + return false +} + +func isApplication(file string) bool { + // tests, test-support packages and everything outside cmd/ and pkg/ are not the application + if strings.HasSuffix(file, "_test.go") || strings.Contains(file, "/testutil/") { + return false + } + if !strings.HasPrefix(file, "cmd/") && !strings.HasPrefix(file, "pkg/") { + return false + } + return !strings.HasSuffix(path.Base(path.Dir(file)), "test") +} diff --git a/hack/check-weak-random/main_test.go b/hack/check-weak-random/main_test.go new file mode 100644 index 000000000..abdf687c2 --- /dev/null +++ b/hack/check-weak-random/main_test.go @@ -0,0 +1,119 @@ +package main + +import ( + "os" + "path/filepath" + "strings" + "testing" +) + +func source(imports ...string) []byte { + var sb strings.Builder + sb.WriteString("package example\n\nimport (\n") + for _, imp := range imports { + sb.WriteString("\t\"" + imp + "\"\n") + } + sb.WriteString(")\n") + return []byte(sb.String()) +} + +func TestViolations(t *testing.T) { + tests := []struct { + name string + file string + imports []string + want int + }{ + {"app math/rand", "pkg/util/middleware/mirror.go", []string{"math/rand"}, 1}, + {"app math/rand/v2", "pkg/util/middleware/mirror.go", []string{"fmt", "math/rand/v2"}, 1}, + {"test math/rand", "pkg/timeseries/dataset/dataset_test.go", []string{"math/rand"}, 1}, + {"tooling math/rand", "hack/load-test/main.go", []string{"math/rand/v2"}, 1}, + {"standalone hrw", "pkg/lb/hrw/hrw.go", []string{"math/rand/v2"}, 0}, + {"standalone p2c", "pkg/lb/p2c/p2c.go", []string{"math/rand/v2"}, 0}, + {"standalone rr", "pkg/lb/rr/rr.go", []string{"math/rand/v2"}, 0}, + {"standalone test helper", "pkg/lb/lbtest/lbtest.go", []string{"math/rand/v2"}, 0}, + {"standalone old rand", "pkg/lb/rr/rr.go", []string{"math/rand"}, 1}, + {"new standalone use", "pkg/lb/rr/new.go", []string{"math/rand/v2"}, 1}, + {"other standalone strategy", "pkg/lb/lc/lc.go", []string{"math/rand/v2"}, 1}, + {"app compat", "pkg/util/middleware/mirror.go", []string{compatPath}, 0}, + {"app weaktest", "pkg/util/middleware/mirror.go", []string{weaktestPath}, 1}, + {"main weak", "cmd/trickster/main.go", []string{"github.com/trickstercache/trickster/v2/pkg/util/weak"}, 0}, + {"test compat", "pkg/util/middleware/mirror_test.go", []string{compatPath}, 1}, + {"test weaktest", "pkg/util/middleware/mirror_test.go", []string{weaktestPath}, 0}, + {"testutil compat", "pkg/testutil/albpool/albpool.go", []string{compatPath}, 1}, + {"testutil weaktest", "pkg/testutil/albpool/albpool.go", []string{weaktestPath}, 0}, + {"helper package weaktest", "pkg/lb/lbtest/lbtest.go", []string{weaktestPath}, 0}, + {"helper package compat", "pkg/lb/lbtest/lbtest.go", []string{compatPath}, 1}, + {"integration weaktest", "integration/main_test.go", []string{weaktestPath}, 0}, + {"integration compat", "integration/harness.go", []string{compatPath}, 1}, + {"weak packages", "pkg/util/weak/compat/compat.go", []string{"math/rand/v2"}, 0}, + {"weak tests", "pkg/util/weak/weak_test.go", []string{compatPath, weaktestPath}, 0}, + {"everything wrong", "pkg/proxy/proxy.go", []string{"math/rand", weaktestPath, "math/rand/v2"}, 3}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + got, err := violations(test.file, source(test.imports...)) + if err != nil { + t.Fatal(err) + } + if len(got) != test.want { + t.Errorf("got %d violations %q, want %d", len(got), got, test.want) + } + }) + } + if _, err := violations("pkg/x.go", []byte("not go")); err == nil { + t.Error("expected a parse error") + } +} + +func TestCheck(t *testing.T) { + root := t.TempDir() + write := func(file string, imports ...string) { + p := filepath.Join(root, filepath.FromSlash(file)) + if err := os.MkdirAll(filepath.Dir(p), 0o755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(p, source(imports...), 0o600); err != nil { + t.Fatal(err) + } + } + write("pkg/app/app.go", compatPath) + write("pkg/app/app_test.go", weaktestPath) + write("pkg/app/testdata/ignored.go", "math/rand") + write("pkg/vendor/ignored.go", "math/rand") + write("pkg/app/README.md") + problems, err := check(root) + if err != nil { + t.Fatal(err) + } + if len(problems) != 0 { + t.Fatalf("unexpected problems: %q", problems) + } + + write("cmd/bad/main.go", weaktestPath) + problems, err = check(root) + if err != nil { + t.Fatal(err) + } + if len(problems) != 1 || !strings.HasPrefix(problems[0], "cmd/bad/main.go:") { + t.Fatalf("unexpected problems: %q", problems) + } + + write("pkg/broken.go") + if err := os.WriteFile(filepath.Join(root, "pkg", "broken.go"), []byte("not go"), 0o600); err != nil { + t.Fatal(err) + } + if _, err := check(root); err == nil { + t.Fatal("expected a parse error") + } +} + +func TestRepositoryConforms(t *testing.T) { + problems, err := check(filepath.Join("..", "..")) + if err != nil { + t.Fatal(err) + } + for _, p := range problems { + t.Error(p) + } +} diff --git a/hack/developer-seed-data.sh b/hack/developer-seed-data.sh index 1abeb5246..fac7e84ad 100755 --- a/hack/developer-seed-data.sh +++ b/hack/developer-seed-data.sh @@ -10,9 +10,8 @@ set -euo pipefail cd "$(dirname "$0")/../docs/developer/environment" # Every trips database service has a one-shot loader service _seed. -# graphite is seeded by its own generator, and prometheus by a backfill of the -# trips metrics that devorigin serves; both are handled separately below. -ALL_TARGETS="clickhouse mysql timescaledb druid prometheus graphite" +# graphite is seeded by its own generator and is handled separately below. +ALL_TARGETS="clickhouse mysql timescaledb greptimedb druid prometheus graphite" read -r -a targets <<< "$(echo "${SEED_TARGET:-$ALL_TARGETS}" | tr ',' ' ')" trips_databases=() @@ -50,6 +49,27 @@ seed_prometheus() { docker compose up -d --no-deps prometheus devorigin } +# developer-start can return while its one-shot loaders are still running. +# All trips loaders share the fixture, even when only one database is reloaded; +# prometheus_seed_generate reads it too, and ALL_TARGETS yields prometheus_seed. +startup_services=() +if [[ ${#trips_databases[@]} -gt 0 || $prometheus -eq 1 ]]; then + startup_services+=(seed_data_generate prometheus_seed_generate) + for db in $ALL_TARGETS; do + if [[ "$db" != graphite ]]; then startup_services+=("${db}_seed"); fi + done +fi +if [[ $graphite -eq 1 ]]; then startup_services+=(graphite_seed); fi +startup_ids=$(docker compose ps -q --status running "${startup_services[@]}") +for id in $startup_ids; do + echo "waiting for startup seeder $id" + status=$(docker wait "$id") + if [[ "$status" != 0 ]]; then + echo "startup seeder $id: FAILED (exit $status)" >&2 + exit 1 + fi +done + names=() pids=() if [[ ${#trips_databases[@]} -gt 0 || $prometheus -eq 1 ]]; then diff --git a/hack/greptimedb-check.sh b/hack/greptimedb-check.sh new file mode 100644 index 000000000..1ca856bc3 --- /dev/null +++ b/hack/greptimedb-check.sh @@ -0,0 +1,58 @@ +#!/bin/sh +# +# Copyright 2026 The Trickster Authors +# 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. + +# Read-only by default; native and PromQL checks own isolated fixture data. +set -eu + +case "${1:-}" in + ""|--promql|--mysql) ;; + *) printf 'Usage: %s [--promql|--mysql]\n' "$0" >&2; exit 2 ;; +esac + +root=$(CDPATH= cd -- "$(dirname -- "$0")/.." && pwd) +cd "$root/integration" +reports=${GREPTIMEDB_REPORT_ROOT:-"$PWD/greptimedb/reports"} +mkdir -p "$reports" +reports=$(CDPATH= cd -- "$reports" && pwd) +GREPTIMEDB_REPORT_DIR=$(mktemp -d "$reports/run-$(date -u +%Y%m%dT%H%M%SZ)-XXXXXX") +export GREPTIMEDB_REPORT_DIR +if [ "${1:-}" = --mysql ]; then + export TRICKSTER_GREPTIMEDB_MYSQL_ACCEPTANCE=1 + status=0 + "${GO:-go}" test -json -count=1 -timeout 10m -run '^TestGreptimeMySQL' . > "$GREPTIMEDB_REPORT_DIR/go-test.jsonl" 2>&1 || status=$? + printf 'MySQL test exit code: %s\nEvidence: %s\n' "$status" "$GREPTIMEDB_REPORT_DIR" + exit "$status" +fi +if [ "${1:-}" = --promql ]; then + export TRICKSTER_GREPTIMEDB_PROMQL_ACCEPTANCE=1 + status=0 + "${GO:-go}" test -json -count=1 -timeout 10m -run '^TestGreptimeDBPrometheus$' . > "$GREPTIMEDB_REPORT_DIR/go-test.jsonl" 2>&1 || status=$? + printf 'PromQL test exit code: %s\nEvidence: %s\n' "$status" "$GREPTIMEDB_REPORT_DIR" + exit "$status" +fi +export TRICKSTER_GREPTIMEDB_ACCEPTANCE=1 + +printf 'Report directory: %s\n' "$GREPTIMEDB_REPORT_DIR" +status=0 +"${GO:-go}" test -json -count=1 -timeout 10m ./greptimedb > "$GREPTIMEDB_REPORT_DIR/go-test.jsonl" 2>&1 || status=$? +printf 'Test exit code: %s\n' "$status" +printf 'Report: %s/report.json\nLog: %s/go-test.jsonl\n' "$GREPTIMEDB_REPORT_DIR" "$GREPTIMEDB_REPORT_DIR" +if [ "${TRICKSTER_GREPTIMEDB_PROXY_ACCEPTANCE:-0}" = 1 ]; then + printf 'Proxy report: %s/proxy-report.json\n' "$GREPTIMEDB_REPORT_DIR" +fi +if [ "${TRICKSTER_GREPTIMEDB_HTTP_ACCEPTANCE:-0}" = 1 ]; then + printf 'HTTP SQL report: %s/http-sql-report.json\n' "$GREPTIMEDB_REPORT_DIR" +fi +exit "$status" diff --git a/hack/greptimeseed/README.md b/hack/greptimeseed/README.md new file mode 100644 index 000000000..7aba691da --- /dev/null +++ b/hack/greptimeseed/README.md @@ -0,0 +1,32 @@ +# greptimeseed + +Loads `hack/seedgen`'s shared gzip TSV files into the developer GreptimeDB +instance through authenticated HTTP SQL. No external Go dependencies or +additional fixture downloads are needed. + +```sh +SEED_PROFILE=small SEED_TARGET=greptimedb make developer-seed-data +cd hack/greptimeseed && go test ./... +``` + +Configuration: `GREPTIMEDB_URL`, `GREPTIMEDB_DATABASE`, +`GREPTIMEDB_SEED_USER`, `GREPTIMEDB_SEED_PASSWORD`, `GREPTIMEDB_SEED_DATA`. +Defaults match the Compose developer environment. This is a destructive +fixture loader: each run drops and recreates `trips` and `trips_15m`. +Do not point it at a production database. + +The loader checks the generated metadata and exact TSV column order, shifts +pickup/dropoff timestamps in UTC, regenerates their dates, and appends +`pickup_epoch` in seconds. Timestamp columns use microsecond precision. +Empty numeric fields become NULL; empty text stays empty text. + +SQL INSERT batches preserve the TimescaleDB fixture's DATE, timestamp and +narrow numeric types. In GreptimeDB v1.2.1, line-protocol string fields +cannot populate existing DATE columns. The table uses low-cardinality +`cab_type` and `vendor_id` tags with append mode so coincident trips survive. +An ambiguous failed INSERT is not retried, since doing so could duplicate +rows. Re-run the entire seed instead. + +After loading, the loader builds a static 15-minute rollup and checks row +count, shifted pickup/dropoff bounds, centering on the seed instant, +date/datetime and epoch agreement, and the rollup's total row count. diff --git a/hack/greptimeseed/go.mod b/hack/greptimeseed/go.mod new file mode 100644 index 000000000..cc5ea13f7 --- /dev/null +++ b/hack/greptimeseed/go.mod @@ -0,0 +1,3 @@ +module github.com/trickstercache/trickster/v2/hack/greptimeseed + +go 1.27 diff --git a/hack/greptimeseed/main.go b/hack/greptimeseed/main.go new file mode 100644 index 000000000..a48922c5d --- /dev/null +++ b/hack/greptimeseed/main.go @@ -0,0 +1,426 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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. + */ + +// Command greptimeseed loads the shared synthetic trips through GreptimeDB's +// HTTP SQL endpoint. SQL preserves DATE, timestamp and narrow numeric types +// that the InfluxDB line protocol cannot write into the existing schema. +package main + +import ( + "bufio" + "compress/gzip" + "encoding/json" + "errors" + "fmt" + "io" + "math" + "net/http" + "net/url" + "os" + "path/filepath" + "strconv" + "strings" + "time" +) + +type column struct { + name string + kind string +} + +// Source order matches hack/seedgen. pickup_epoch is appended during loading. +var columns = []column{ + {"trip_id", "BIGINT"}, + {"vendor_id", "STRING"}, + {"pickup_date", "DATE"}, + {"pickup_datetime", "TIMESTAMP(6)"}, + {"dropoff_date", "DATE"}, + {"dropoff_datetime", "TIMESTAMP(6)"}, + {"store_and_fwd_flag", "SMALLINT"}, + {"rate_code_id", "SMALLINT"}, + {"pickup_longitude", "DOUBLE"}, + {"pickup_latitude", "DOUBLE"}, + {"dropoff_longitude", "DOUBLE"}, + {"dropoff_latitude", "DOUBLE"}, + {"passenger_count", "SMALLINT"}, + {"trip_distance", "DOUBLE"}, + {"fare_amount", "FLOAT"}, + {"extra", "FLOAT"}, + {"transit_tax", "FLOAT"}, + {"tip_amount", "FLOAT"}, + {"tolls_amount", "FLOAT"}, + {"ehail_fee", "FLOAT"}, + {"improvement_surcharge", "FLOAT"}, + {"total_amount", "FLOAT"}, + {"payment_type", "STRING"}, + {"trip_type", "SMALLINT"}, + {"pickup", "STRING"}, + {"dropoff", "STRING"}, + {"cab_type", "STRING"}, + {"pickup_zone_gid", "INT"}, + {"pickup_tract_label", "FLOAT"}, + {"pickup_borough_code", "SMALLINT"}, + {"pickup_borough_name", "STRING"}, + {"pickup_tract_code", "STRING"}, + {"pickup_district_class", "STRING"}, + {"pickup_neighborhood_code", "STRING"}, + {"pickup_neighborhood_name", "STRING"}, + {"pickup_ward", "INT"}, + {"dropoff_zone_gid", "INT"}, + {"dropoff_tract_label", "FLOAT"}, + {"dropoff_borough_code", "SMALLINT"}, + {"dropoff_borough_name", "STRING"}, + {"dropoff_tract_code", "STRING"}, + {"dropoff_district_class", "STRING"}, + {"dropoff_neighborhood_code", "STRING"}, + {"dropoff_neighborhood_name", "STRING"}, + {"dropoff_ward", "INT"}, +} + +var metadataKeys = []string{ + "SOURCE_ROWS", "SOURCE_PICKUP_MIN_EPOCH", "SOURCE_PICKUP_MAX_EPOCH", + "SOURCE_DROPOFF_MIN_EPOCH", "SOURCE_DROPOFF_MAX_EPOCH", "SEED_EPOCH", "SHIFT_SECONDS", +} + +type seeder struct { + baseURL, database, user, password, dataDir string + client *http.Client + log io.Writer +} + +type sqlOutput struct { + AffectedRows *int64 `json:"affectedrows"` + Records *struct { + Rows [][]json.Number `json:"rows"` + } `json:"records"` +} + +func main() { + s := &seeder{ + baseURL: envOr("GREPTIMEDB_URL", "http://greptimedb:4000"), //nolint:revive // developer-environment default + database: envOr("GREPTIMEDB_DATABASE", "public"), + user: envOr("GREPTIMEDB_SEED_USER", "seeder"), + password: envOr("GREPTIMEDB_SEED_PASSWORD", "trickster-dev-seed"), + dataDir: envOr("GREPTIMEDB_SEED_DATA", "/seed-data"), + client: &http.Client{Timeout: 2 * time.Minute}, + log: os.Stdout, + } + if err := s.run(); err != nil { + _, _ = fmt.Fprintln(os.Stderr, "greptime seed:", err) + os.Exit(1) + } +} + +func envOr(key, fallback string) string { + if v := os.Getenv(key); v != "" { + return v + } + return fallback +} + +func (s *seeder) sql(query string) (sqlOutput, error) { + form := url.Values{"db": {s.database}, "sql": {query}} + req, err := http.NewRequest(http.MethodPost, strings.TrimRight(s.baseURL, "/")+"/v1/sql", strings.NewReader(form.Encode())) + if err != nil { + return sqlOutput{}, err + } + req.Header.Set("Content-Type", "application/x-www-form-urlencoded") + req.SetBasicAuth(s.user, s.password) + resp, err := s.client.Do(req) + if err != nil { + return sqlOutput{}, err + } + defer resp.Body.Close() + var doc struct { + Output []sqlOutput `json:"output"` + Code int `json:"code"` + Error string `json:"error"` + } + if err := json.NewDecoder(io.LimitReader(resp.Body, 1<<20)).Decode(&doc); err != nil { + return sqlOutput{}, fmt.Errorf("SQL HTTP %d: %w", resp.StatusCode, err) + } + if resp.StatusCode != http.StatusOK || doc.Code != 0 || doc.Error != "" { + return sqlOutput{}, fmt.Errorf("SQL HTTP %d code %d: %s", resp.StatusCode, doc.Code, doc.Error) + } + if len(doc.Output) != 1 { + return sqlOutput{}, fmt.Errorf("expected one SQL result, got %d", len(doc.Output)) + } + return doc.Output[0], nil +} + +func readMetadata(path string) (map[string]int64, error) { + raw, err := os.ReadFile(path) + if err != nil { + return nil, err + } + m := make(map[string]int64) + for line := range strings.SplitSeq(string(raw), "\n") { + line = strings.TrimSpace(line) + if line == "" { + continue + } + key, value, ok := strings.Cut(line, "=") + if !ok { + return nil, fmt.Errorf("invalid metadata line %q", line) + } + n, err := strconv.ParseInt(value, 10, 64) + if err != nil { + return nil, fmt.Errorf("invalid %s: %w", key, err) + } + if _, exists := m[key]; exists { + return nil, fmt.Errorf("duplicate metadata key %s", key) + } + m[key] = n + } + for _, key := range metadataKeys { + if _, ok := m[key]; !ok { + return nil, fmt.Errorf("missing metadata key %s", key) + } + } + if m["SOURCE_ROWS"] <= 0 { + return nil, errors.New("SOURCE_ROWS must be positive") + } + for _, kind := range []string{"PICKUP", "DROPOFF"} { + minTime, maxTime := m["SOURCE_"+kind+"_MIN_EPOCH"], m["SOURCE_"+kind+"_MAX_EPOCH"] + if minTime > maxTime { + return nil, fmt.Errorf("reversed %s bounds", kind) + } + for _, v := range []int64{minTime, maxTime} { + if _, err := shiftEpoch(v, m["SHIFT_SECONDS"]); err != nil { + return nil, err + } + } + } + return m, nil +} + +func shiftEpoch(epoch, shift int64) (time.Time, error) { + if shift > 0 && epoch > math.MaxInt64-shift || shift < 0 && epoch < math.MinInt64-shift { + return time.Time{}, errors.New("timestamp shift overflows") + } + t := time.Unix(epoch+shift, 0).UTC() + if t.Year() < 1 || t.Year() > 9999 { + return time.Time{}, errors.New("timestamp outside supported calendar range") + } + return t, nil +} + +func sourceNames() []string { + names := make([]string, len(columns)) + for i, c := range columns { + names[i] = c.name + } + return names +} + +func createTableSQL() string { + defs := make([]string, 0, len(columns)+3) + for _, c := range columns { + def := `"` + c.name + `" ` + c.kind + if c.name == "pickup_datetime" { + def += " NOT NULL TIME INDEX" + } else if c.name == "trip_id" || c.name == "vendor_id" || c.name == "cab_type" { + def += " NOT NULL" + } + defs = append(defs, def) + } + defs = append(defs, "pickup_epoch BIGINT NOT NULL", "PRIMARY KEY (cab_type, vendor_id)") + return "CREATE TABLE trips (" + strings.Join(defs, ",") + ") WITH (append_mode='true')" +} + +func quote(s string) string { return "'" + strings.ReplaceAll(s, "'", "''") + "'" } + +func rowSQL(fields []string, shift int64) (string, error) { + if len(fields) != len(columns) { + return "", fmt.Errorf("expected %d columns, got %d", len(columns), len(fields)) + } + values := make([]string, len(fields)+1) + var pickup time.Time + for i, c := range columns { + value := fields[i] + switch c.kind { + case "STRING": + values[i] = quote(value) + case "DATE": + // Date fields are regenerated from the shifted datetime next to them. + continue + case "TIMESTAMP(6)": + if value == "" && c.name == "dropoff_datetime" { + values[i-1], values[i] = "NULL", "NULL" + continue + } + t, err := time.Parse("2006-01-02 15:04:05", value) + if err != nil { + return "", fmt.Errorf("%s: %w", c.name, err) + } + t, err = shiftEpoch(t.Unix(), shift) + if err != nil { + return "", err + } + values[i-1] = quote(t.Format(time.DateOnly)) + values[i] = quote(t.Format(time.RFC3339)) + if c.name == "pickup_datetime" { + pickup = t + } + default: + if value == "" { + values[i] = "NULL" + continue + } + if c.kind == "FLOAT" || c.kind == "DOUBLE" { + n, err := strconv.ParseFloat(value, 64) + if err != nil || math.IsNaN(n) || math.IsInf(n, 0) { + return "", fmt.Errorf("invalid %s: %q", c.name, value) + } + values[i] = strconv.FormatFloat(n, 'g', -1, 64) + } else { + bits := 64 + if c.kind == "SMALLINT" { + bits = 16 + } else if c.kind == "INT" { + bits = 32 + } + n, err := strconv.ParseInt(value, 10, bits) + if err != nil { + return "", fmt.Errorf("%s: %w", c.name, err) + } + values[i] = strconv.FormatInt(n, 10) + } + } + } + values[len(fields)] = strconv.FormatInt(pickup.Unix(), 10) + return "(" + strings.Join(values, ",") + ")", nil +} + +func (s *seeder) loadFile(name string, shift int64) (int64, error) { + f, err := os.Open(filepath.Join(s.dataDir, name)) + if err != nil { + return 0, err + } + defer f.Close() + zr, err := gzip.NewReader(f) + if err != nil { + return 0, err + } + defer zr.Close() + scan := bufio.NewScanner(zr) + scan.Buffer(make([]byte, 64<<10), 1<<20) + if !scan.Scan() || scan.Text() != strings.Join(sourceNames(), "\t") { + return 0, fmt.Errorf("%s: missing or mismatched TSV header", name) + } + const batchSize = 1000 + batch := make([]string, 0, batchSize) + var rows int64 + flush := func() error { + if len(batch) == 0 { + return nil + } + // Do not retry an INSERT: append mode would duplicate a committed batch + // whose response was lost. A failed seed must restart from DROP TABLE. + out, err := s.sql(`INSERT INTO trips ("` + strings.Join(sourceNames(), `","`) + `",pickup_epoch) VALUES ` + strings.Join(batch, ",")) + if err != nil { + return err + } + if out.AffectedRows == nil || *out.AffectedRows != int64(len(batch)) { + return errors.New("INSERT affected-row count does not match batch") + } + rows += int64(len(batch)) + batch = batch[:0] + return nil + } + for scan.Scan() { + row, err := rowSQL(strings.Split(scan.Text(), "\t"), shift) + if err != nil { + return rows, fmt.Errorf("%s row %d: %w", name, rows+int64(len(batch))+1, err) + } + batch = append(batch, row) + if len(batch) == batchSize { + if err := flush(); err != nil { + return rows, err + } + } + } + if err := scan.Err(); err != nil { + return rows, fmt.Errorf("%s: %w", name, err) + } + if err := flush(); err != nil { + return rows, err + } + return rows, nil +} + +func (s *seeder) run() error { + m, err := readMetadata(filepath.Join(s.dataDir, "seed-window.env")) + if err != nil { + return err + } + for _, q := range []string{"DROP TABLE IF EXISTS trips_15m", "DROP TABLE IF EXISTS trips", createTableSQL()} { + if _, err := s.sql(q); err != nil { + return err + } + } + for _, name := range []string{"trips_1.gz", "trips_2.gz"} { + if _, err := s.loadFile(name, m["SHIFT_SECONDS"]); err != nil { + return err + } + } + for _, q := range []string{ + `CREATE TABLE trips_15m ("bucket" TIMESTAMP(6) TIME INDEX, cab_type STRING, trips BIGINT, total_amount_sum DOUBLE, PRIMARY KEY (cab_type))`, + `INSERT INTO trips_15m SELECT date_bin(INTERVAL '15 minutes', pickup_datetime) AS "bucket", cab_type, count(*), sum(total_amount) FROM trips GROUP BY 1, 2`, + } { + if _, err := s.sql(q); err != nil { + return err + } + } + return s.validate(m) +} + +func (s *seeder) validate(m map[string]int64) error { + out, err := s.sql(`SELECT count(*), + CAST(extract(epoch FROM min(pickup_datetime)) AS BIGINT), + CAST(extract(epoch FROM max(pickup_datetime)) AS BIGINT), + CAST(extract(epoch FROM min(dropoff_datetime)) AS BIGINT), + CAST(extract(epoch FROM max(dropoff_datetime)) AS BIGINT), + count(*) FILTER (WHERE pickup_date <> CAST(pickup_datetime AS DATE)), + count(*) FILTER (WHERE dropoff_datetime IS NOT NULL AND dropoff_date IS DISTINCT FROM CAST(dropoff_datetime AS DATE)), + count(*) FILTER (WHERE pickup_epoch <> CAST(extract(epoch FROM pickup_datetime) AS BIGINT)), + (SELECT CAST(sum(trips) AS BIGINT) FROM trips_15m) + FROM trips`) + if err != nil { + return err + } + want := []int64{ + m["SOURCE_ROWS"], + m["SOURCE_PICKUP_MIN_EPOCH"] + m["SHIFT_SECONDS"], m["SOURCE_PICKUP_MAX_EPOCH"] + m["SHIFT_SECONDS"], + m["SOURCE_DROPOFF_MIN_EPOCH"] + m["SHIFT_SECONDS"], m["SOURCE_DROPOFF_MAX_EPOCH"] + m["SHIFT_SECONDS"], + 0, 0, 0, m["SOURCE_ROWS"], + } + if out.Records == nil || len(out.Records.Rows) != 1 || len(out.Records.Rows[0]) != len(want) { + return errors.New("invalid seed validation result") + } + for i, expected := range want { + got, err := out.Records.Rows[0][i].Int64() + if err != nil || got != expected { + return fmt.Errorf("seed validation column %d: got %s, want %d", i, out.Records.Rows[0][i], expected) + } + } + midpoint := want[1] + (want[2]-want[1])/2 + if midpoint != m["SEED_EPOCH"] { + return errors.New("seed window is not centered on SEED_EPOCH") + } + _, _ = fmt.Fprintf(s.log, "seed complete: %d rows, pickup window %d..%d\n", want[0], want[1], want[2]) + return nil +} diff --git a/hack/greptimeseed/main_test.go b/hack/greptimeseed/main_test.go new file mode 100644 index 000000000..d47912644 --- /dev/null +++ b/hack/greptimeseed/main_test.go @@ -0,0 +1,305 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 ( + "bytes" + "compress/gzip" + "fmt" + "io" + "math" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "testing" +) + +const metadata = "SOURCE_ROWS=2\nSOURCE_PICKUP_MIN_EPOCH=1704067200\nSOURCE_PICKUP_MAX_EPOCH=1704067202\n" + + "SOURCE_DROPOFF_MIN_EPOCH=1704067260\nSOURCE_DROPOFF_MAX_EPOCH=1704067262\nSEED_EPOCH=1704153601\nSHIFT_SECONDS=86400\n" + +func fixture() []string { + r := make([]string, len(columns)) + for i, c := range columns { + switch c.kind { + case "DATE": + r[i] = "2024-01-01" + case "TIMESTAMP(6)": + r[i] = "2024-01-01 00:00:00" + case "STRING": + r[i] = "" + default: + r[i] = "1" + } + } + r[1], r[26], r[5] = "O'Brien\\fleet", "yellow", "2024-01-01 00:01:00" + return r +} + +func TestRowSQL(t *testing.T) { + r := fixture() + r[12] = "" + got, err := rowSQL(r, -1) + if err != nil { + t.Fatal(err) + } + for _, want := range []string{"'O''Brien\\fleet'", "'2023-12-31','2023-12-31T23:59:59Z'", "'2024-01-01','2024-01-01T00:00:59Z'", ",NULL,", "1704067199)"} { + if !strings.Contains(got, want) { + t.Errorf("missing %q in %s", want, got) + } + } + r[5] = "" + got, err = rowSQL(r, 0) + if err != nil || !strings.Contains(got, ",NULL,NULL,") { + t.Fatalf("nullable dropoff: %s, %v", got, err) + } + for _, tc := range []struct { + index int + value string + }{{3, ""}, {3, "not a timestamp"}, {12, "32768"}, {0, "1;DROP TABLE trips"}, {14, "NaN"}, {14, "Inf"}, {14, "1e999"}} { + r := fixture() + r[tc.index] = tc.value + if _, err := rowSQL(r, 0); err == nil { + t.Errorf("accepted column %d value %q", tc.index, tc.value) + } + } + if _, err := rowSQL(r[:3], 0); err == nil { + t.Fatal("accepted short row") + } + for _, shift := range []int64{math.MaxInt64, math.MinInt64} { + if _, err := rowSQL(fixture(), shift); err == nil { + t.Errorf("accepted shift %d", shift) + } + } +} + +func TestMetadata(t *testing.T) { + for _, tc := range []struct { + name, text string + valid bool + }{ + {"valid", metadata, true}, + {"CRLF", strings.ReplaceAll(metadata, "\n", "\r\n"), true}, + {"missing", strings.ReplaceAll(metadata, "SHIFT_SECONDS=86400\n", ""), false}, + {"zero rows", strings.ReplaceAll(metadata, "SOURCE_ROWS=2", "SOURCE_ROWS=0"), false}, + {"expression", strings.ReplaceAll(metadata, "SHIFT_SECONDS=86400", "SHIFT_SECONDS=$(date)"), false}, + {"duplicate", metadata + "SHIFT_SECONDS=1\n", false}, + {"invalid line", metadata + "broken\n", false}, + {"reversed", strings.ReplaceAll(metadata, "MIN_EPOCH=1704067200", "MIN_EPOCH=1704067300"), false}, + {"overflow", strings.ReplaceAll(metadata, "SHIFT_SECONDS=86400", "SHIFT_SECONDS=9223372036854775807"), false}, + } { + t.Run(tc.name, func(t *testing.T) { + path := filepath.Join(t.TempDir(), "seed-window.env") + if err := os.WriteFile(path, []byte(tc.text), 0o600); err != nil { + t.Fatal(err) + } + _, err := readMetadata(path) + if (err == nil) != tc.valid { + t.Fatalf("valid=%v, err=%v", tc.valid, err) + } + }) + } +} + +func seedFile(t *testing.T, rows int, header string) string { + t.Helper() + dir := t.TempDir() + var b bytes.Buffer + w := gzip.NewWriter(&b) + _, _ = fmt.Fprintln(w, header) + for range rows { + _, _ = fmt.Fprintln(w, strings.Join(fixture(), "\t")) + } + if err := w.Close(); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(dir, "trips.gz"), b.Bytes(), 0o600); err != nil { + t.Fatal(err) + } + return dir +} + +func TestLoadFile(t *testing.T) { + for _, size := range []int{1, 1000, 1001} { + t.Run(fmt.Sprint(size), func(t *testing.T) { + requests, inserted := 0, 0 + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + requests++ + if err := r.ParseForm(); err != nil { + t.Error(err) + } + u, p, ok := r.BasicAuth() + if !ok || u != "seeder" || p != "fixture-password" || r.Form.Get("db") != "public" || r.URL.Path != "/v1/sql" { + t.Errorf("incorrect SQL request: %s", r.URL) + } + query := r.Form.Get("sql") + n := strings.Count(query, "),(") + 1 + inserted += n + _, _ = fmt.Fprintf(w, `{"output":[{"affectedrows":%d}]}`, n) + })) + defer srv.Close() + s := &seeder{ + baseURL: srv.URL, database: "public", user: "seeder", password: "fixture-password", client: srv.Client(), + dataDir: seedFile(t, size, strings.Join(sourceNames(), "\t")), log: io.Discard, + } + rows, err := s.loadFile("trips.gz", 86400) + if err != nil || rows != int64(size) || inserted != size || requests != (size+999)/1000 { + t.Fatalf("rows=%d inserted=%d requests=%d err=%v", rows, inserted, requests, err) + } + }) + } +} + +func TestSQLFailuresNotRetried(t *testing.T) { + for _, tc := range []struct { + status int + body string + }{ + {500, `{"error":"unavailable"}`}, + {200, `{"code":1000,"error":"failed"}`}, + {200, `{"output":[]}`}, + {200, `not json`}, + {200, `{"output":[{"affectedrows":0}]}`}, + {200, `{"output":[{}]}`}, + } { + t.Run(tc.body, func(t *testing.T) { + calls := 0 + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + calls++ + w.WriteHeader(tc.status) + _, _ = io.WriteString(w, tc.body) + })) + defer srv.Close() + s := &seeder{baseURL: srv.URL, client: srv.Client(), dataDir: seedFile(t, 1, strings.Join(sourceNames(), "\t"))} + if _, err := s.loadFile("trips.gz", 0); err == nil || calls != 1 { + t.Fatalf("calls=%d err=%v", calls, err) + } + }) + } +} + +func TestBadSeedFile(t *testing.T) { + for _, corrupt := range []bool{false, true} { + dir := seedFile(t, 1, "wrong header") + if corrupt { + dir = seedFile(t, 1, strings.Join(sourceNames(), "\t")) + path := filepath.Join(dir, "trips.gz") + b, err := os.ReadFile(path) + if err != nil { + t.Fatal(err) + } + b[len(b)-8] ^= 1 // corrupt the gzip checksum + if err := os.WriteFile(path, b, 0o600); err != nil { + t.Fatal(err) + } + } + s := &seeder{dataDir: dir} + if _, err := s.loadFile("trips.gz", 0); err == nil { + t.Fatalf("accepted corrupt=%v", corrupt) + } + } +} + +const facts = `{"output":[{"records":{"rows":[[2,1704153600,1704153602,1704153660,1704153662,0,0,0,2]]}}]}` + +func TestValidate(t *testing.T) { + for _, tc := range []struct { + name, response string + valid bool + }{ + {"valid", facts, true}, + {"lost duplicate", strings.Replace(facts, "[[2,", "[[1,", 1), false}, + {"wrong bounds", strings.Replace(facts, "1704153600", "1704153601", 1), false}, + {"date mismatch", strings.Replace(facts, ",0,0,0,2", ",1,0,0,2", 1), false}, + {"incomplete rollup", strings.Replace(facts, ",0,0,0,2", ",0,0,0,1", 1), false}, + {"NULL bound", strings.Replace(facts, "1704153600", "null", 1), false}, + {"missing records", `{"output":[{}]}`, false}, + {"short row", `{"output":[{"records":{"rows":[[2]]}}]}`, false}, + } { + t.Run(tc.name, func(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + _, _ = io.WriteString(w, tc.response) + })) + defer srv.Close() + s := &seeder{baseURL: srv.URL, client: srv.Client(), log: io.Discard} + path := filepath.Join(t.TempDir(), "metadata") + if err := os.WriteFile(path, []byte(metadata), 0o600); err != nil { + t.Fatal(err) + } + m, err := readMetadata(path) + if err != nil { + t.Fatal(err) + } + if err := s.validate(m); (err == nil) != tc.valid { + t.Fatalf("valid=%v, err=%v", tc.valid, err) + } + m["SEED_EPOCH"]++ + if err := s.validate(m); err == nil { + t.Fatal("accepted uncentered window") + } + }) + } +} + +func TestRunRecreatesTables(t *testing.T) { + dir := seedFile(t, 1, strings.Join(sourceNames(), "\t")) + b, err := os.ReadFile(filepath.Join(dir, "trips.gz")) + if err != nil { + t.Fatal(err) + } + for _, name := range []string{"trips_1.gz", "trips_2.gz"} { + if err := os.WriteFile(filepath.Join(dir, name), b, 0o600); err != nil { + t.Fatal(err) + } + } + if err := os.WriteFile(filepath.Join(dir, "seed-window.env"), []byte(metadata), 0o600); err != nil { + t.Fatal(err) + } + var queries []string + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if err := r.ParseForm(); err != nil { + t.Error(err) + } + q := r.Form.Get("sql") + queries = append(queries, q) + switch { + case strings.HasPrefix(q, "SELECT count(*)"): + _, _ = io.WriteString(w, facts) + case strings.HasPrefix(q, "INSERT INTO trips ("): + _, _ = io.WriteString(w, `{"output":[{"affectedrows":1}]}`) + default: + _, _ = io.WriteString(w, `{"output":[{"affectedrows":0}]}`) + } + })) + defer srv.Close() + s := &seeder{baseURL: srv.URL, client: srv.Client(), log: io.Discard, dataDir: dir} + for range 2 { + queries = nil + if err := s.run(); err != nil { + t.Fatal(err) + } + if len(queries) != 8 || queries[0] != "DROP TABLE IF EXISTS trips_15m" || queries[1] != "DROP TABLE IF EXISTS trips" { + t.Fatalf("unexpected seed lifecycle: %v", queries) + } + for _, clause := range []string{`"pickup_datetime" TIMESTAMP(6) NOT NULL TIME INDEX`, "PRIMARY KEY (cab_type, vendor_id)", "append_mode='true'"} { + if !strings.Contains(queries[2], clause) { + t.Errorf("schema missing %s", clause) + } + } + } +} diff --git a/integration/Makefile b/integration/Makefile index 862914c8e..584afeb32 100644 --- a/integration/Makefile +++ b/integration/Makefile @@ -24,15 +24,18 @@ export QUIC_GO_DISABLE_RECEIVE_BUFFER_WARNING = 1 test: $(GO) test -v --failfast + $(GO) test -v ./greptimedb @../hack/filter-coverprofile.sh .integration.coverprofile data-race-test: $(GO) test -v --failfast --race + $(GO) test -v --race ./greptimedb cover: $(GO) test -v --failfast -coverpkg=../... -coverprofile=.integration.coverprofile + $(GO) test -v ./greptimedb @../hack/filter-coverprofile.sh .integration.coverprofile $(GO) tool cover -func=.integration.coverprofile | tail -1 cover-html: cover - $(GO) tool cover -html=.integration.coverprofile \ No newline at end of file + $(GO) tool cover -html=.integration.coverprofile diff --git a/integration/greptimedb/README.md b/integration/greptimedb/README.md new file mode 100644 index 000000000..435306e4d --- /dev/null +++ b/integration/greptimedb/README.md @@ -0,0 +1,241 @@ +# GreptimeDB Environment Acceptance + +This read-only suite checks the running developer environment and, optionally, +the pgwire and HTTP SQL caches. It does not start containers, seed data, change database +contents, or approve a development phase. Use only the isolated developer +environment: the default credentials and plaintext connections are for it. + +## Run + +After starting and seeding the developer environment, from the repository root: + +```sh +make developer-greptimedb-check +``` + +To check Trickster, start it with the developer configuration: + +```sh +TRICKSTER_GREPTIMEDB_PROXY_ACCEPTANCE=1 make developer-greptimedb-check +``` + +This additionally writes `proxy-report.json`: pgwire simple/extended/error +recovery, invalid credentials, HTTP SQL/PromQL, provider metrics, and all +twelve SQL dashboard panels compared directly with the same GreptimeDB origin. +Each panel runs on its initial range, repeats it, then requests an overlapping +range. Counter deltas verify the expected object/delta path, a repeated hit, +and a partial hit (or an already warmed hit) for overlapping delta ranges. +SQL rewrite failures fail the check even when a fallback returns correct data. +Every field, row and nanosecond is exact, including panel 6 percentages; the +cross-engine rounding allowance below does not apply to proxy comparisons. +Only HTTP SQL execution duration is excluded from its data comparison. +Provide `GREPTIMEDB_PASSTHROUGH_PG_ADDR` and `GREPTIMEDB_PASSTHROUGH_SQL_UID` +for a second listener/datasource without an authenticator to verify passthrough +as well. Without both, that mode is explicitly UNVERIFIED, not a passed check. + +After restarting Prometheus, allow at least four minutes of healthy scrapes +before running the sample-grid check. A startup window with missing samples is +still a failed run; retain that report separately from a later qualified run. + +To include HTTP SQL cache comparisons: + +```sh +TRICKSTER_GREPTIMEDB_HTTP_ACCEPTANCE=1 make developer-greptimedb-check +``` + +This writes `http-sql-report.json` and the paired origin/proxy responses for +GET and POST requests. It checks miss, partial hit, hit, empty results, exact +integers above 2^53, timezone identity, alternate formats, and unaligned ranges. +Typed schema, row order and numeric values must agree with the origin; only +execution duration is excluded. Unaligned raw-time bounds and inclusive upper +bounds are checked against explicit complete-bucket reference SQL; the proxy +receives the original unaligned SQL and must use delta caching, including on +a cold request. Both submitted and reference queries are retained. +Authenticated GET object responses are not stored unless the origin marks them +shareable; a repeat miss in that case is the expected HTTP cache policy. + +## PromQL Provider And Merge Acceptance + +The PromQL suite is separate from the read-only checks above: + +```sh +sh hack/greptimedb-check.sh --promql +``` + +Run from the repository root on Linux with the isolated developer GreptimeDB +running. It uses the public developer `seeder` account to create three uniquely +named `trickster_promql_*` databases, seeds a small exact-value fixture, starts +its own Trickster daemon on reserved ports, then stops that daemon and drops +only the databases it created. It does not reseed or drop `public` tables. +Do not run it against a production or shared database. After a killed process, +inspect its retained config to identify any fixture databases needing cleanup. + +Every origin and proxy response is retained in a new `promql-*` directory along +with the generated config and the parent `go-test.jsonl`. The suite checks +GET/POST miss/hit/partial/hit, alignment of shifted and 500ms grids, database/header/lookback +isolation, metadata, empty/error responses, URL-over-form precedence, and +two-member merges against the complete dataset. Numeric sample spelling may +differ (`2` versus `2.0`); numeric values and point timestamps must be equal to +the explicitly step-aligned origin request. Passthrough cases compare original +requests without alignment. +Series and metadata sets have no defined order. The shared merger's exact +range-sort advisory is expected separately; other warnings are not discarded. + +`count_values` currently loses its numeric grouping label in GreptimeDB's wire +response. The suite preserves that observed origin result through passthrough +and requires TSM to reject it instead of returning an incorrect aggregate. +This is an explicit upstream limitation, not a successful merge case. + +## MySQL Provider Acceptance + +```sh +sh hack/greptimedb-check.sh --mysql +``` + +This opt-in suite creates a uniquely named `trickster_mysql_*` table using the +public developer seeder, queries it as a read-only user, and drops that table +after its loopback listener has stopped. Use only an isolated developer database. +The queries compare raw typed results, not rounded floating-point conversions. +They check object and delta misses/hits, an extended-range partial hit, NULL and +case-sensitive label ordering, timestamps with nine fractional digits, integers +above 2^53, error recovery, failed SET, partial buckets, and stateful bypass. +Counters must show the intended cache path; equal results alone are insufficient. +The script retains a unique `go-test.jsonl` even on failure. An unavailable origin +fails explicitly requested acceptance instead of silently skipping it. + +## Evidence + +The script uses the integration module's Go version and dependencies. Each run +creates a new directory under `integration/greptimedb/reports/` containing: + +- `report.json`: individual PASS/FAIL/UNSUPPORTED/UNVERIFIED checks, observed versions, + timestamps, input hashes and PostgreSQL type/parameter observations. +- `go-test.jsonl`: Go's machine-readable test log, including failures. +- `panel-*.json`, `seed-*.json`, `promql-*.json`: query responses for comparison. + +A nonzero exit means the automatic checks failed (or did not finish). A missing +or empty report is not success. Earlier results are never overwritten. Passing +automatic checks does not make UNVERIFIED items pass or finish issue #1150. +UNSUPPORTED records an observed upstream limitation, not a skipped successful +test. Its corresponding capability probe must pass before that label is used. + +GreptimeDB v1.2.1 fails Grafana's SQL health check and the comment-only +PostgreSQL query check. The developer environment now pins the official +`nightly-20260923-e91faa9df` image containing the upstream repair. Keep the +old-version failed runs; do not disable assertions or describe a source-built +repair or an official nightly as a tested stable release. + +## Assertions + +- The Docker lifecycle script's unit tests run with a fake Docker executable, + not a live daemon. They require Bash on Linux and check that active startup + seeders finish before shared fixture generation or table reloads. A failed + wait must prevent all mutations. The real delete/start/reseed lifecycle is + still a separate, destructive reviewer check. +- Grafana and all three required direct datasource health checks must succeed. +- GreptimeDB and TimescaleDB counts and pickup bounds must match the shared + `seed-window.env` metadata. +- The six main dashboard panels are read from the checked-in JSON. Both + datasources use the same fixed 48-hour window ending at UTC midnight on + `SEED_EPOCH`'s date and a five-minute interval. Compare every field name, type, + label, value, nanosecond remainder and row order. Plugin-specific metadata + such as executed SQL and display configuration is intentionally excluded. +- Panel 6's `card_use_rate` alone permits adjacent finite float64 values in + the range 0-100. PostgreSQL numeric division and DataFusion floating division + can differ by one rounding step (observed for 37 non-cash trips out of 99). + The report records the number of such cells and retains both raw responses. + Larger differences, all other fields, integers and timestamps remain exact. +- Every requested Grafana query ref must have a successful, nonempty result. + Embedded query errors, missing refs, ragged columns and all-null data fail + even when the HTTP response is 200. JSON integers retain their precision. +- `up{job="prometheus"}` must contain the same nine healthy samples on a 15-second + grid through both Prometheus datasources. Its separately recorded window + ends one minute before the run starts. Allow at least four minutes of healthy + scraping and remote write before starting the suite; it does not retry an + empty response until it happens to pass. +- Direct PostgreSQL simple/extended queries and timestamp 0/3/6/9, DATE, + integer and float OIDs/text are checked. Startup observations come from + **pgconn, not captured Grafana traffic**. The observed PG text renderer uses + six fractional digits even for `TIMESTAMP(9)`; this is not evidence of + end-to-end nanosecond preservation. A provider must match the upstream wire + contract instead of inferring it from the storage type. +- Direct HTTP SQL and MySQL ping/query must preserve a BIGINT above 2^53. +- PostgreSQL and MySQL bucket queries are compared with raw `pickup_epoch` + rows grouped independently in Go. The suite checks `extract`, `date_part`, + `date_bin` with interval/compact widths, and `::interval` casts. Both the + original window and a window shortened by 37 seconds at each end must use + the same epoch-aligned five-minute grid and the correct counts. Each SQL + form and its result are retained, including failures. +- MySQL requires an explicit interval unit: `INTERVAL '5' MINUTE`. Its rejection + of the PostgreSQL spelling `INTERVAL '5 minutes'` must be a server syntax + error, not a timeout or lost connection. Literal/identifier quoting is also + checked separately for each protocol. +- PostgreSQL startup currently supplies zero cancellation credentials, matching + the upstream no-op cancellation handler. `BEGIN`/`ROLLBACK` are compatibility + stubs with an explicit no-transactions warning. The suite checks the observed + `T`, `T`, `E`, `I`, `I` status sequence and query values before/after an error; + it does not claim storage transactions or isolation. Keep cancellation and + transaction support disabled until upstream behavior is reevaluated. + +## Grafana Wire Evidence + +The typed query also runs through Grafana's actual bundled PostgreSQL plugin. +`grafana-type-query.json` is an API response, not a protocol capture. A separate +capture must observe Grafana connecting directly to GreptimeDB without a shim. +Capture a fresh connection and health check, then run this suite. Check startup +parameters, all server ParameterStatus messages, query mode, RowDescription +OIDs/formats, and the typed row's original text. Do not substitute the pgconn +probe for this evidence or infer nanosecond preservation from API timestamps. + +Scope the capture to the isolated Grafana peer and PostgreSQL port. Export +selected protocol fields only; do not retain passwords, authentication data, +cancellation secrets or raw PCAP files. Save the capture and its independent +review alongside the report. The suite's own UNVERIFIED entry intentionally +remains unchanged because it neither performs nor validates that capture. + +## Overrides + +| Variable | Default | Purpose | +| --- | --- | --- | +| `GO` | `go` | Go executable | +| `GRAFANA_URL` | `http://127.0.0.1:3000` | Anonymous developer Grafana API | +| `GREPTIMEDB_PG_ADDR` | `127.0.0.1:4003` | PostgreSQL address | +| `GREPTIMEDB_MYSQL_ADDR` | `127.0.0.1:4002` | MySQL address | +| `GREPTIMEDB_HTTP_URL` | `http://127.0.0.1:4000` | HTTP SQL base URL | +| `GREPTIMEDB_PROXY_HTTP_URL` | `http://127.0.0.1:8480/greptimedb1` | Trickster HTTP base URL | +| `GREPTIMEDB_DATABASE` | `public` | Database | +| `GREPTIMEDB_USER` | `grafana_ro` | Read-only developer account | +| `GREPTIMEDB_PASSWORD` | `trickster-dev-grafana` | Developer password; not written to the report | +| `GREPTIMEDB_SQL_UID` | `ds_greptimedb_direct` | Grafana SQL datasource UID | +| `GREPTIMEDB_PROM_UID` | `ds_greptimedb_prom_direct` | Grafana Prometheus datasource UID | +| `GREPTIMEDB_REPORT_ROOT` | `integration/greptimedb/reports` | Parent for unique run directories | +| `GREPTIMEDB_BUILD_NOTE` | unset | Operator-supplied image/source identification, not independently attested | + +For remote Docker validation, run the same command inside the remote checkout +or Go container and set the addresses to that environment. Never expose a +developer database publicly to make these checks reachable. + +Ordinary integration test targets run the negative/unit cases without reaching +external services. To run them alone: + +```sh +cd integration +go test -count=1 ./greptimedb +go test -race -count=1 ./greptimedb +``` + +## Human Review + +1. Read `report.json` and any FAIL details alongside `go-test.jsonl`. Check the + observed server version and build note before comparing separate runs. +2. Open the GreptimeDB and TimescaleDB dashboards at exactly `sql_from` through + `sql_to`, using the direct datasources. Inspect all six data panels on desktop + and mobile. An HTTP/API assertion is not a rendering assertion. +3. Open the existing Prometheus dashboard using `greptimedb-prom-direct`. + This is distinct from the GreptimeDB dashboard's Trickster performance + panels, which may remain empty until the later provider/cache phases. +4. Review the UNVERIFIED list against the issue checklist. Actual Grafana wire + capture and complete reseed coverage need separate evidence. Check the + capability probes before accepting UNSUPPORTED entries. +5. Record the human decision and evidence separately. This suite does not + approve its own results, establish correctness beyond the tested cases, or submit a PR. diff --git a/integration/greptimedb/acceptance_test.go b/integration/greptimedb/acceptance_test.go new file mode 100644 index 000000000..9ab0508cc --- /dev/null +++ b/integration/greptimedb/acceptance_test.go @@ -0,0 +1,397 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 greptimedb_test + +import ( + "crypto/sha256" + "encoding/json" + "fmt" + "net/http" + "net/url" + "os" + "path/filepath" + "runtime" + "strconv" + "strings" + "testing" + "time" +) + +const environment = "../../docs/developer/environment/docker-compose-data" + +type dashboard struct { + Panels []struct { + ID int `json:"id"` + Title string `json:"title"` + Targets []map[string]any `json:"targets"` + } `json:"panels"` +} + +type check struct { + Name string `json:"name"` + Status string `json:"status"` + Detail string `json:"detail,omitempty"` +} + +type report struct { + StartedAt time.Time `json:"started_at"` + GoVersion string `json:"go_version"` + BuildNote string `json:"operator_build_note,omitempty"` + From time.Time `json:"sql_from"` + To time.Time `json:"sql_to"` + Hashes map[string]string `json:"input_sha256"` + Facts map[string]any `json:"observed_facts"` + Checks []check `json:"checks"` +} + +func envOr(name, fallback string) string { + if v := os.Getenv(name); v != "" { + return v + } + return fallback +} + +func readSeed(raw []byte) (map[string]int64, error) { + keys := []string{ + "SOURCE_ROWS", "SOURCE_PICKUP_MIN_EPOCH", "SOURCE_PICKUP_MAX_EPOCH", + "SOURCE_DROPOFF_MIN_EPOCH", "SOURCE_DROPOFF_MAX_EPOCH", "SEED_EPOCH", "SHIFT_SECONDS", + } + m := make(map[string]int64, len(keys)) + for line := range strings.SplitSeq(string(raw), "\n") { + line = strings.TrimSpace(line) + if line == "" { + continue + } + key, value, ok := strings.Cut(line, "=") + if !ok { + return nil, fmt.Errorf("invalid seed metadata line %q", line) + } + if _, ok := m[key]; ok { + return nil, fmt.Errorf("duplicate seed key %s", key) + } + n, err := strconv.ParseInt(value, 10, 64) + if err != nil { + return nil, fmt.Errorf("invalid seed value for %s: %w", key, err) + } + m[key] = n + } + for _, key := range keys { + if _, ok := m[key]; !ok { + return nil, fmt.Errorf("missing seed key %s", key) + } + } + if len(m) != len(keys) || m["SOURCE_ROWS"] <= 0 || m["SEED_EPOCH"] <= 0 || + m["SOURCE_PICKUP_MIN_EPOCH"] > m["SOURCE_PICKUP_MAX_EPOCH"] || + m["SOURCE_DROPOFF_MIN_EPOCH"] > m["SOURCE_DROPOFF_MAX_EPOCH"] { + return nil, fmt.Errorf("invalid seed metadata") + } + return m, nil +} + +func writeJSON(path string, v any) error { + b, err := json.MarshalIndent(v, "", " ") + if err != nil { + return err + } + return os.WriteFile(path, append(b, '\n'), 0o600) +} + +func (g grafanaClient) query(targets []map[string]any, uid, kind string, from, to time.Time, step time.Duration) (queryResponse, error) { + queries := make([]map[string]any, len(targets)) + refs := make([]string, len(targets)) + for i, target := range targets { + q := make(map[string]any, len(target)+3) + for key, value := range target { + q[key] = value + } + q["datasource"] = map[string]string{"type": kind, "uid": uid} + q["intervalMs"] = step.Milliseconds() + q["maxDataPoints"] = int(to.Sub(from)/step) + 1 + queries[i] = q + refs[i], _ = q["refId"].(string) + } + var doc queryResponse + err := g.request(http.MethodPost, "/api/ds/query", map[string]any{ + "queries": queries, "from": strconv.FormatInt(from.UnixMilli(), 10), "to": strconv.FormatInt(to.UnixMilli(), 10), + }, &doc) + if err == nil { + err = validateResponse(doc, refs) + } + return doc, err +} + +func TestDirectEnvironment(t *testing.T) { + if os.Getenv("TRICKSTER_GREPTIMEDB_ACCEPTANCE") != "1" { + t.Skip("set TRICKSTER_GREPTIMEDB_ACCEPTANCE=1 for the read-only live suite") + } + out := os.Getenv("GREPTIMEDB_REPORT_DIR") + if out == "" { + t.Fatal("GREPTIMEDB_REPORT_DIR is required; use make developer-greptimedb-check") + } + if err := os.MkdirAll(out, 0o700); err != nil { + t.Fatal(err) + } + // Never overwrite an earlier run, including its failures. + f, err := os.OpenFile(filepath.Join(out, "report.json"), os.O_CREATE|os.O_EXCL|os.O_WRONLY, 0o600) + if err != nil { + t.Fatal(err) + } + if err := f.Close(); err != nil { + t.Fatal(err) + } + r := report{ + StartedAt: time.Now().UTC(), GoVersion: runtime.Version(), + BuildNote: os.Getenv("GREPTIMEDB_BUILD_NOTE"), Hashes: map[string]string{}, Facts: map[string]any{}, + } + t.Cleanup(func() { + for _, capability := range []struct{ probe, name, detail string }{ + {"pg_cancel_capability", "pg_cancel_support", "Zero cancellation credentials; upstream uses a no-op cancellation handler"}, + {"pg_transaction_compatibility", "pg_transaction_support", "BEGIN/ROLLBACK are compatibility stubs; upstream warns that transactions are unsupported"}, + } { + c := check{capability.name, "UNVERIFIED", "Capability probe did not pass"} + for _, probe := range r.Checks { + if probe.Name == capability.probe && probe.Status == "PASS" { + c.Status, c.Detail = "UNSUPPORTED", capability.detail + } + } + r.Checks = append(r.Checks, c) + } + for _, name := range []string{ + "actual Grafana PostgreSQL wire capture (not the pgconn probe)", + "desktop/mobile visual inspection of every data panel", + "existing Prometheus dashboard rendered against the GreptimeDB PromQL datasource", + "full-profile reseed and existing seed targets", + "populated GreptimeDB Trickster performance panels (later provider/cache phases)", + "Trickster provider, cache behavior and complete issue #1150 acceptance", + } { + r.Checks = append(r.Checks, check{name, "UNVERIFIED", "Outside this read-only direct-environment suite"}) + } + if err := writeJSON(filepath.Join(out, "report.json"), r); err != nil { + t.Error(err) + } + t.Logf("review report: %s", filepath.Join(out, "report.json")) + }) + run := func(name string, fn func() error) { + t.Run(name, func(t *testing.T) { + c := check{Name: name, Status: "PASS"} + if err := fn(); err != nil { + c.Status, c.Detail = "FAIL", err.Error() + t.Error(err) + } + r.Checks = append(r.Checks, c) + }) + } + read := func(name, path string, dst any) error { + b, err := os.ReadFile(path) + if err != nil { + return err + } + r.Hashes[name] = fmt.Sprintf("%x", sha256.Sum256(b)) + return json.Unmarshal(b, dst) + } + var seed map[string]int64 + var greptime, timescale dashboard + var setupErr error + run("inputs", func() error { + raw, err := os.ReadFile(filepath.Join(environment, "seed-data", "seed-window.env")) + if err == nil { + seed, err = readSeed(raw) + r.Hashes["seed-window.env"] = fmt.Sprintf("%x", sha256.Sum256(raw)) + } + if err == nil { + err = read("greptimedb-dashboard", filepath.Join(environment, "dashboards", "trickster-greptimedb.json"), &greptime) + } + if err == nil { + err = read("timescaledb-dashboard", filepath.Join(environment, "dashboards", "trickster-timescaledb.json"), ×cale) + } + setupErr = err + return err + }) + if setupErr != nil { + return + } + r.To = time.Unix(seed["SEED_EPOCH"], 0).UTC().Truncate(24 * time.Hour) + r.From = r.To.Add(-48 * time.Hour) + r.Facts["seed_metadata"] = seed + g := grafanaClient{strings.TrimRight(envOr("GRAFANA_URL", "http://127.0.0.1:3000"), "/"), &http.Client{Timeout: 30 * time.Second}} + const sqlKind = "grafana-postgresql-datasource" + gtUID := envOr("GREPTIMEDB_SQL_UID", "ds_greptimedb_direct") + tsUID := "ds_timescaledb_direct" + promUID := envOr("GREPTIMEDB_PROM_UID", "ds_greptimedb_prom_direct") + run("grafana_version", func() error { + var health struct{ Version, Database string } + if err := g.request(http.MethodGet, "/api/health", nil, &health); err != nil { + return err + } + r.Facts["grafana_version"] = health.Version + if health.Version == "" || health.Database != "ok" { + return fmt.Errorf("Grafana is not healthy: %+v", health) + } + return nil + }) + for _, uid := range []string{gtUID, tsUID, promUID} { + run("health_"+uid, func() error { + var health struct{ Status, Message string } + if err := g.request(http.MethodGet, "/api/datasources/uid/"+url.PathEscape(uid)+"/health", nil, &health); err != nil { + return err + } + if health.Status != "OK" { + return fmt.Errorf("datasource status %q: %s", health.Status, health.Message) + } + return nil + }) + } + querySQL := func(uid, sql string) (queryResponse, error) { + return g.query([]map[string]any{{"refId": "A", "format": "table", "rawQuery": true, "rawSql": sql}}, uid, sqlKind, r.From, r.To, 5*time.Minute) + } + run("grafana_pg_type_query", func() error { + doc, err := querySQL(gtUID, pgTypeQuery) + if err != nil { + return err + } + r.Facts["grafana_pg_type_frame_not_raw_wire"] = doc.Results["A"].Frames + frames := doc.Results["A"].Frames + if len(frames) != 1 || len(frames[0].Schema.Fields) != 10 || len(frames[0].Data.Values[0]) != 1 { + return fmt.Errorf("unexpected Grafana typed-query frame") + } + return writeJSON(filepath.Join(out, "grafana-type-query.json"), doc) + }) + for _, uid := range []string{gtUID, tsUID} { + run("version_"+uid, func() error { + doc, err := querySQL(uid, "SELECT version() AS version") + if err == nil { + r.Facts["version_"+uid] = doc.Results["A"].Frames + } + return err + }) + } + run("seed_count_and_bounds", func() error { + sql := "SELECT COUNT(*) AS rows, MIN(pickup_epoch) AS first, MAX(pickup_epoch) AS last FROM trips" + want := []any{ + json.Number(strconv.FormatInt(seed["SOURCE_ROWS"], 10)), + json.Number(strconv.FormatInt(seed["SOURCE_PICKUP_MIN_EPOCH"]+seed["SHIFT_SECONDS"], 10)), + json.Number(strconv.FormatInt(seed["SOURCE_PICKUP_MAX_EPOCH"]+seed["SHIFT_SECONDS"], 10)), + } + for _, uid := range []string{gtUID, tsUID} { + doc, err := querySQL(uid, sql) + if err != nil { + return err + } + if err := writeJSON(filepath.Join(out, "seed-"+uid+".json"), doc); err != nil { + return err + } + frames := doc.Results["A"].Frames + if len(frames) != 1 || len(frames[0].Data.Values) != len(want) { + return fmt.Errorf("%s: unexpected seed result shape", uid) + } + for i, value := range want { + if len(frames[0].Data.Values[i]) != 1 || frames[0].Data.Values[i][0] != value { + return fmt.Errorf("%s: seed field %d = %v, want %v", uid, i, frames[0].Data.Values[i], value) + } + } + } + return nil + }) + for id := 1; id <= 6; id++ { + run(fmt.Sprintf("panel_%d", id), func() error { + var docs []queryResponse + for i, dashboard := range []dashboard{timescale, greptime} { + var targets []map[string]any + for _, p := range dashboard.Panels { + if p.ID == id { + targets = p.Targets + } + } + if len(targets) == 0 { + return fmt.Errorf("missing panel %d targets", id) + } + uid := []string{tsUID, gtUID}[i] + doc, err := g.query(targets, uid, sqlKind, r.From, r.To, 5*time.Minute) + if writeErr := writeJSON(filepath.Join(out, fmt.Sprintf("panel-%d-%s.json", id, uid)), doc); writeErr != nil { + return writeErr + } + if err != nil { + return fmt.Errorf("%s: %w", uid, err) + } + docs = append(docs, doc) + } + if id == 6 { + rounded, err := compareFrames(docs[0], docs[1], true) + r.Facts["panel_6_percentage_rounding"] = map[string]any{ + "differing_cells": rounded, "maximum_float64_steps": 1, + "field": "card_use_rate", "other_fields": "exact", + } + return err + } + return compareResponses(docs[0], docs[1]) + }) + } + run("promql_sample_grid", func() error { + // Exclude the newest samples so remote-write lag cannot change the window. + to := r.StartedAt.Add(-time.Minute).Truncate(15 * time.Second) + from := to.Add(-2 * time.Minute) + r.Facts["promql_window"] = map[string]any{"from": from, "to": to, "step_seconds": 15} + targets := []map[string]any{{"refId": "A", "expr": `up{job="prometheus"}`, "range": true, "instant": false, "interval": "15s"}} + var docs []queryResponse + for _, uid := range []string{"ds_prom_direct", promUID} { + doc, err := g.query(targets, uid, "prometheus", from, to, 15*time.Second) + if writeErr := writeJSON(filepath.Join(out, "promql-"+uid+".json"), doc); writeErr != nil { + return writeErr + } + if err != nil { + return err + } + frames := doc.Results["A"].Frames + if len(frames) != 1 || len(frames[0].Data.Values) != 2 || len(frames[0].Data.Values[0]) != 9 { + return fmt.Errorf("%s: expected one complete nine-sample series", uid) + } + for i, value := range frames[0].Data.Values[1] { + wantTime := json.Number(strconv.FormatInt(from.Add(time.Duration(i)*15*time.Second).UnixMilli(), 10)) + if value != json.Number("1") || frames[0].Data.Values[0][i] != wantTime { + return fmt.Errorf("%s: incomplete or unhealthy scrape at sample %d", uid, i) + } + } + docs = append(docs, doc) + } + return compareResponses(docs[0], docs[1]) + }) + protocolChecks(run, &r) +} + +func TestReadSeed(t *testing.T) { + valid := "SOURCE_ROWS=2\nSOURCE_PICKUP_MIN_EPOCH=1\nSOURCE_PICKUP_MAX_EPOCH=2\nSOURCE_DROPOFF_MIN_EPOCH=2\nSOURCE_DROPOFF_MAX_EPOCH=3\nSEED_EPOCH=1800000000\nSHIFT_SECONDS=1799999999\n" + for _, tt := range []struct { + name, raw string + ok bool + }{ + {"valid", valid, true}, + {"crlf", strings.ReplaceAll(valid, "\n", "\r\n"), true}, + {"missing", strings.Replace(valid, "SEED_EPOCH=1800000000\n", "", 1), false}, + {"duplicate", valid + "SEED_EPOCH=1\n", false}, + {"empty_rows", strings.Replace(valid, "SOURCE_ROWS=2", "SOURCE_ROWS=0", 1), false}, + {"shell_value", strings.Replace(valid, "SOURCE_ROWS=2", "SOURCE_ROWS=$(date)", 1), false}, + {"unknown_key", valid + "EXTRA=1\n", false}, + {"reversed_bounds", strings.Replace(valid, "SOURCE_PICKUP_MAX_EPOCH=2", "SOURCE_PICKUP_MAX_EPOCH=0", 1), false}, + } { + t.Run(tt.name, func(t *testing.T) { + _, err := readSeed([]byte(tt.raw)) + if (err == nil) != tt.ok { + t.Fatalf("error = %v, want success %v", err, tt.ok) + } + }) + } +} diff --git a/integration/greptimedb/dialect_test.go b/integration/greptimedb/dialect_test.go new file mode 100644 index 000000000..5604c56ee --- /dev/null +++ b/integration/greptimedb/dialect_test.go @@ -0,0 +1,330 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 greptimedb_test + +import ( + "context" + "database/sql" + "errors" + "fmt" + "reflect" + "strconv" + "strings" + "testing" + "time" + + "github.com/go-sql-driver/mysql" + "github.com/jackc/pgx/v5/pgconn" +) + +type textQuery func(string) ([][]string, error) + +func pgTextQuery(ctx context.Context, c *pgconn.PgConn) textQuery { + return func(query string) ([][]string, error) { + results, err := c.Exec(ctx, query).ReadAll() + if err != nil { + return nil, err + } + if len(results) != 1 { + return nil, fmt.Errorf("expected one PostgreSQL result, got %d", len(results)) + } + rows := make([][]string, len(results[0].Rows)) + for i, row := range results[0].Rows { + rows[i] = make([]string, len(row)) + for j, v := range row { + if v == nil { + return nil, fmt.Errorf("unexpected NULL at row %d column %d", i, j) + } + rows[i][j] = string(v) + } + } + return rows, nil + } +} + +func mysqlTextQuery(ctx context.Context, db *sql.DB) textQuery { + return func(query string) ([][]string, error) { + rows, err := db.QueryContext(ctx, query) + if err != nil { + return nil, err + } + defer rows.Close() + columns, err := rows.Columns() + if err != nil { + return nil, err + } + var result [][]string + for rows.Next() { + values := make([]sql.NullString, len(columns)) + dest := make([]any, len(columns)) + for i := range values { + dest[i] = &values[i] + } + if err := rows.Scan(dest...); err != nil { + return nil, err + } + row := make([]string, len(values)) + for i, value := range values { + if !value.Valid { + return nil, fmt.Errorf("unexpected NULL at column %d", i) + } + row[i] = value.String + } + result = append(result, row) + } + return result, rows.Err() + } +} + +func floorBucket(epoch int64) int64 { + const width = int64(300) + q := epoch / width + if epoch%width < 0 { + q-- + } + return q * width +} + +func bucketOracle(rows [][]string) (map[int64]int64, error) { + if len(rows) == 0 { + return nil, fmt.Errorf("empty raw timestamp fixture") + } + want := make(map[int64]int64) + for _, row := range rows { + if len(row) != 1 { + return nil, fmt.Errorf("expected one raw timestamp column") + } + epoch, err := strconv.ParseInt(row[0], 10, 64) + if err != nil { + return nil, err + } + want[floorBucket(epoch)]++ + } + return want, nil +} + +func checkBuckets(rows [][]string, want map[int64]int64) error { + if len(rows) == 0 || len(want) == 0 { + return fmt.Errorf("empty bucket comparison") + } + got := make(map[int64]int64, len(rows)) + var previous int64 + for i, row := range rows { + if len(row) != 2 { + return fmt.Errorf("expected bucket and count, got %v", row) + } + epoch, err := strconv.ParseInt(row[0], 10, 64) + if err != nil { + return err + } + count, err := strconv.ParseInt(row[1], 10, 64) + if err != nil { + return err + } + if epoch%300 != 0 || count <= 0 || (i > 0 && epoch <= previous) { + return fmt.Errorf("unaligned, unordered, duplicate or empty bucket: %v", row) + } + got[epoch], previous = count, epoch + } + if !reflect.DeepEqual(got, want) { + return fmt.Errorf("bucket counts differ from independently grouped raw rows") + } + return nil +} + +func checkBucketSQL(query textQuery, r *report, protocol string) error { + forms := []struct { + name, expression string + mysqlSyntaxError bool + }{ + {"extract", "CAST(floor(extract(epoch FROM pickup_datetime) / 300) * 300 AS BIGINT)", false}, + {"date_part", "CAST(floor(date_part('epoch', pickup_datetime) / 300) * 300 AS BIGINT)", false}, + {"date_bin_interval", "CAST(extract(epoch FROM date_bin(INTERVAL '5 minutes', pickup_datetime)) AS BIGINT)", true}, + {"date_bin_unit", "CAST(extract(epoch FROM date_bin(INTERVAL '5' MINUTE, pickup_datetime)) AS BIGINT)", false}, + {"date_bin_compact", "CAST(extract(epoch FROM date_bin('5m', pickup_datetime)) AS BIGINT)", false}, + {"interval_cast", "CAST(extract(epoch FROM date_bin('5 minutes'::interval, pickup_datetime)) AS BIGINT)", false}, + } + var evidence []map[string]any + var failures []error + defer func() { r.Facts[protocol+"_bucket_cases"] = evidence }() + for _, shift := range []time.Duration{0, 37 * time.Second} { + from, to := r.From.Add(shift), r.To.Add(-shift) + where := fmt.Sprintf(" WHERE pickup_datetime >= '%s' AND pickup_datetime < '%s'", from.Format(time.RFC3339), to.Format(time.RFC3339)) + raw, err := query("SELECT pickup_epoch FROM trips" + where) + if err != nil { + return err + } + want, err := bucketOracle(raw) + if err != nil { + return err + } + for _, form := range forms { + sql := "SELECT " + form.expression + " AS bucket_epoch, count(*) AS n FROM trips" + where + " GROUP BY bucket_epoch ORDER BY bucket_epoch" + rows, err := query(sql) + item := map[string]any{"form": form.name, "sql": sql, "from": from, "to": to, "source_rows": len(raw), "buckets": rows} + evidence = append(evidence, item) + if form.mysqlSyntaxError && protocol == "mysql" { + // MySQL requires INTERVAL expr unit; an arbitrary query failure is not evidence of rejection. + var syntaxErr *mysql.MySQLError + if errors.As(err, &syntaxErr) && syntaxErr.Number == 1149 && string(syntaxErr.SQLState[:]) == "42000" { + item["expected_syntax_error"] = syntaxErr.Error() + continue + } + err = fmt.Errorf("expected MySQL interval syntax error, got %v", err) + } + if err == nil { + err = checkBuckets(rows, want) + } + if err != nil { + item["error"] = err.Error() + failures = append(failures, fmt.Errorf("%s, shift %s: %w", form.name, shift, err)) + } + } + } + return errors.Join(failures...) +} + +func checkQuoting(query textQuery, r *report, protocol string) error { + var evidence []map[string]any + defer func() { r.Facts[protocol+"_quoting_cases"] = evidence }() + for _, tc := range []struct { + sql, want string + mysqlOnly bool + }{ + {"SELECT 'literal' AS answer", "literal", false}, + {`SELECT "literal" AS answer`, "literal", true}, + {"SELECT 42 AS `answer`", "42", true}, + } { + rows, err := query(tc.sql) + item := map[string]any{"sql": tc.sql, "rows": rows} + if err != nil { + item["error"] = err.Error() + } + evidence = append(evidence, item) + if tc.mysqlOnly && protocol == "pg" { + if err == nil { + return fmt.Errorf("PostgreSQL unexpectedly accepts MySQL quoting: %s", tc.sql) + } + continue + } + if err != nil { + return err + } + if !reflect.DeepEqual(rows, [][]string{{tc.want}}) { + return fmt.Errorf("%s: got %v, want %s", tc.sql, rows, tc.want) + } + } + return nil +} + +func TestBucketOracle(t *testing.T) { + for _, tc := range []struct{ epoch, want int64 }{{-301, -600}, {-300, -300}, {-1, -300}, {0, 0}, {299, 0}, {300, 300}, {301, 300}} { + t.Run(strconv.FormatInt(tc.epoch, 10), func(t *testing.T) { + if got := floorBucket(tc.epoch); got != tc.want { + t.Fatalf("bucket = %d, want %d", got, tc.want) + } + }) + } + want, err := bucketOracle([][]string{{"0"}, {"0"}, {"299"}, {"300"}}) + if err != nil || !reflect.DeepEqual(want, map[int64]int64{0: 3, 300: 1}) { + t.Fatalf("duplicates must count independently: %v, %v", want, err) + } + for _, rows := range [][][]string{nil, {{"1", "2"}}, {{"1.5"}}, {{"invalid"}}} { + if _, err := bucketOracle(rows); err == nil { + t.Fatalf("accepted invalid raw rows %v", rows) + } + } +} + +func TestCheckBuckets(t *testing.T) { + want := map[int64]int64{0: 3, 300: 1} + for _, tc := range []struct { + name string + rows [][]string + ok bool + }{ + {"valid", [][]string{{"0", "3"}, {"300", "1"}}, true}, + {"shifted_grid", [][]string{{"37", "3"}, {"337", "1"}}, false}, + {"wrong_count", [][]string{{"0", "2"}, {"300", "1"}}, false}, + {"missing_bucket", [][]string{{"0", "3"}}, false}, + {"reversed_order", [][]string{{"300", "1"}, {"0", "3"}}, false}, + {"duplicate_bucket", [][]string{{"0", "1"}, {"0", "2"}, {"300", "1"}}, false}, + {"empty", nil, false}, + {"malformed", [][]string{{"0"}}, false}, + {"non_integer", [][]string{{"0.5", "3"}, {"300", "1"}}, false}, + {"empty_count", [][]string{{"0", "0"}, {"300", "1"}}, false}, + } { + t.Run(tc.name, func(t *testing.T) { + err := checkBuckets(tc.rows, want) + if (err == nil) != tc.ok { + t.Fatalf("error = %v, want success %v", err, tc.ok) + } + }) + } +} + +func TestCheckBucketSQL(t *testing.T) { + for _, tc := range []struct { + name, protocol, failure string + ok bool + }{ + {"postgres", "pg", "", true}, + {"mysql", "mysql", "", true}, + {"mysql_missing_rejection", "mysql", "accept_interval", false}, + {"mysql_transport_error", "mysql", "transport", false}, + {"mysql_wrong_error", "mysql", "wrong_error", false}, + {"postgres_query_error", "pg", "query", false}, + {"wrong_counts", "pg", "counts", false}, + {"missing_raw_rows", "pg", "raw", false}, + } { + t.Run(tc.name, func(t *testing.T) { + r := report{From: time.Unix(0, 0), To: time.Unix(900, 0), Facts: map[string]any{}} + query := func(sql string) ([][]string, error) { + if strings.HasPrefix(sql, "SELECT pickup_epoch") { + if tc.failure == "raw" { + return nil, nil + } + return [][]string{{"100"}, {"400"}}, nil + } + if tc.failure == "query" { + return nil, fmt.Errorf("query failed") + } + if tc.protocol == "mysql" && strings.Contains(sql, "INTERVAL '5 minutes'") { + switch tc.failure { + case "transport": + return nil, fmt.Errorf("connection lost") + case "wrong_error": + return nil, &mysql.MySQLError{Number: 1045} + case "accept_interval": + default: + return nil, &mysql.MySQLError{Number: 1149, SQLState: [5]byte{'4', '2', '0', '0', '0'}} + } + } + if tc.failure == "counts" { + return [][]string{{"0", "2"}, {"300", "1"}}, nil + } + return [][]string{{"0", "1"}, {"300", "1"}}, nil + } + if err := checkBucketSQL(query, &r, tc.protocol); (err == nil) != tc.ok { + t.Fatalf("error = %v, want success %v", err, tc.ok) + } + if tc.failure != "raw" && len(r.Facts[tc.protocol+"_bucket_cases"].([]map[string]any)) != 12 { + t.Fatal("must retain every SQL form in both windows, including after failures") + } + }) + } +} diff --git a/integration/greptimedb/grafana_test.go b/integration/greptimedb/grafana_test.go new file mode 100644 index 000000000..679187bf0 --- /dev/null +++ b/integration/greptimedb/grafana_test.go @@ -0,0 +1,372 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 greptimedb_test + +import ( + "bytes" + "encoding/json" + "fmt" + "io" + "math" + "net/http" + "net/http/httptest" + "reflect" + "strings" + "testing" + "time" +) + +type field struct { + Name string `json:"name"` + Type string `json:"type"` + Labels map[string]string `json:"labels,omitempty"` +} + +type frame struct { + Schema struct { + Name string `json:"name"` + RefID string `json:"refId"` + Fields []field `json:"fields"` + } `json:"schema"` + Data struct { + Values [][]any `json:"values"` + Nanos [][]any `json:"nanos,omitempty"` + } `json:"data"` +} + +type queryResult struct { + Status int `json:"status"` + Error string `json:"error"` + Frames []frame `json:"frames"` +} + +type queryResponse struct { + Results map[string]queryResult `json:"results"` +} + +type grafanaClient struct { + baseURL string + client *http.Client +} + +func (g grafanaClient) request(method, path string, body any, dst any) error { + var data []byte + var err error + if body != nil { + data, err = json.Marshal(body) + if err != nil { + return err + } + } + req, err := http.NewRequest(method, g.baseURL+path, bytes.NewReader(data)) + if err != nil { + return err + } + req.Header.Set("Content-Type", "application/json") + resp, err := g.client.Do(req) + if err != nil { + return err + } + defer resp.Body.Close() + const maxBody = 16 << 20 + raw, err := io.ReadAll(io.LimitReader(resp.Body, maxBody+1)) + if err != nil { + return err + } + if len(raw) > maxBody { + return fmt.Errorf("%s: response exceeds %d bytes", path, maxBody) + } + if resp.StatusCode != http.StatusOK { + return fmt.Errorf("%s: HTTP %d: %.512s", path, resp.StatusCode, raw) + } + d := json.NewDecoder(bytes.NewReader(raw)) + d.UseNumber() + if err := d.Decode(dst); err != nil { + return fmt.Errorf("%s: %w", path, err) + } + if err := d.Decode(new(any)); err != io.EOF { + return fmt.Errorf("%s: trailing JSON data", path) + } + return nil +} + +func validateResponse(doc queryResponse, refs []string) error { + if len(refs) == 0 || len(doc.Results) != len(refs) { + return fmt.Errorf("result refs do not match requested refs %v", refs) + } + seen := make(map[string]bool, len(refs)) + for _, ref := range refs { + if ref == "" || seen[ref] { + return fmt.Errorf("empty or duplicate ref %q", ref) + } + seen[ref] = true + r, ok := doc.Results[ref] + if !ok || r.Status != http.StatusOK || r.Error != "" || len(r.Frames) == 0 { + return fmt.Errorf("ref %s: missing or unsuccessful result (status %d, error %q)", ref, r.Status, r.Error) + } + for i, f := range r.Frames { + fields, columns := len(f.Schema.Fields), len(f.Data.Values) + if f.Schema.RefID != ref || fields == 0 || fields != columns { + return fmt.Errorf("ref %s frame %d: invalid schema/column count", ref, i) + } + rows := len(f.Data.Values[0]) + if rows == 0 { + return fmt.Errorf("ref %s frame %d: empty rows", ref, i) + } + hasValue := false + for j, col := range f.Data.Values { + if len(col) != rows { + return fmt.Errorf("ref %s frame %d column %d: inconsistent row count", ref, i, j) + } + for _, v := range col { + hasValue = hasValue || (v != nil && f.Schema.Fields[j].Type != "time") + } + } + if !hasValue { + return fmt.Errorf("ref %s frame %d: no non-null data values", ref, i) + } + if len(f.Data.Nanos) != 0 && len(f.Data.Nanos) != columns { + return fmt.Errorf("ref %s frame %d: invalid nanosecond column count", ref, i) + } + for _, col := range f.Data.Nanos { + if col != nil && len(col) != rows { + return fmt.Errorf("ref %s frame %d: inconsistent nanosecond row count", ref, i) + } + } + } + } + return nil +} + +func compareResponses(left, right queryResponse) error { + _, err := compareFrames(left, right, false) + return err +} + +func compareFrames(left, right queryResponse, allowPercentageRounding bool) (int, error) { + rounded := 0 + if len(left.Results) != len(right.Results) { + return 0, fmt.Errorf("Grafana result counts differ") + } + for ref, l := range left.Results { + r, ok := right.Results[ref] + if !ok || len(l.Frames) != len(r.Frames) { + return 0, fmt.Errorf("ref %s: frame counts differ", ref) + } + for i, f := range l.Frames { + if !reflect.DeepEqual(f.Schema, r.Frames[i].Schema) { + return 0, fmt.Errorf("ref %s frame %d: schemas differ", ref, i) + } + other := r.Frames[i].Data + if !reflect.DeepEqual(f.Data.Nanos, other.Nanos) || len(f.Data.Values) != len(other.Values) { + return 0, fmt.Errorf("ref %s frame %d: nanoseconds or column counts differ", ref, i) + } + for col, values := range f.Data.Values { + if len(values) != len(other.Values[col]) { + return 0, fmt.Errorf("ref %s frame %d column %d: row counts differ", ref, i, col) + } + for row, value := range values { + actual := other.Values[col][row] + if reflect.DeepEqual(value, actual) { + continue + } + if allowPercentageRounding && col < len(f.Schema.Fields) && + f.Schema.Fields[col].Name == "card_use_rate" && f.Schema.Fields[col].Type == "number" && + adjacentPercentages(value, actual) { + rounded++ + continue + } + return 0, fmt.Errorf("ref %s frame %d column %d row %d: value %v differs from %v", ref, i, col, row, value, actual) + } + } + } + } + return rounded, nil +} + +// PostgreSQL numeric division and DataFusion floating division can round a +// percentage to adjacent float64 values. Never apply this to counts or times. +func adjacentPercentages(left, right any) bool { + a, ok := left.(json.Number) + if !ok { + return false + } + b, ok := right.(json.Number) + if !ok { + return false + } + x, errX := a.Float64() + y, errY := b.Float64() + if errX != nil || errY != nil || !(x >= 0 && x <= 100 && y >= 0 && y <= 100) { + return false + } + return x == y || math.Nextafter(x, y) == y +} + +const validResponse = `{"results":{"A":{"status":200,"frames":[{"schema":{"refId":"A","fields":[{"name":"time","type":"time"},{"name":"count","type":"number","labels":{"job":"prom"}}]},"data":{"values":[[1000,2000],[9007199254740993,2]],"nanos":[[1,2],null]}}]}}}` + +func TestPercentageRounding(t *testing.T) { + const exact, adjacent = "37.37373737373737", "37.37373737373738" + raw := strings.ReplaceAll(strings.Replace(validResponse, `"name":"count"`, `"name":"card_use_rate"`, 1), "9007199254740993", exact) + left := decodeFixture(t, raw) + right := decodeFixture(t, strings.Replace(raw, exact, adjacent, 1)) + if err := compareResponses(left, right); err == nil { + t.Fatal("strict comparison accepted different numbers") + } + if n, err := compareFrames(left, right, true); err != nil || n != 1 { + t.Fatalf("rounded cells = %d, error = %v", n, err) + } + for _, tt := range []struct{ name, old, replacement string }{ + {"two_steps", exact, "37.373737373737385"}, + {"null", exact, "null"}, + {"timestamp", "[1000,2000]", "[1001,2000]"}, + {"nanoseconds", `"nanos":[[1,2],null]`, `"nanos":[[1,3],null]`}, + {"label", `"job":"prom"`, `"job":"other"`}, + {"field_type", `"type":"number"`, `"type":"string"`}, + } { + t.Run(tt.name, func(t *testing.T) { + if _, err := compareFrames(left, decodeFixture(t, strings.Replace(raw, tt.old, tt.replacement, 1)), true); err == nil { + t.Fatal("accepted a material difference") + } + }) + } + if _, err := compareFrames(decodeFixture(t, validResponse), decodeFixture(t, strings.Replace(validResponse, "9007199254740993", "9007199254740992", 1)), true); err == nil { + t.Fatal("percentage tolerance lost integer precision") + } +} + +func TestAdjacentPercentages(t *testing.T) { + for _, tt := range []struct{ name, left, right string }{ + {"negative", "-1", "-1.0000000000000002"}, + {"above_100", "101", "101.00000000000001"}, + {"nan", "NaN", "NaN"}, + {"infinity", "+Inf", "+Inf"}, + {"overflow", "1e9999", "1e9999"}, + {"malformed", "not-a-number", "1"}, + } { + t.Run(tt.name, func(t *testing.T) { + if adjacentPercentages(json.Number(tt.left), json.Number(tt.right)) { + t.Fatal("accepted invalid percentage") + } + }) + } + if adjacentPercentages(nil, json.Number("1")) || adjacentPercentages(json.Number("1"), "1") { + t.Fatal("accepted nonnumeric type") + } +} + +func decodeFixture(t *testing.T, raw string) queryResponse { + t.Helper() + var doc queryResponse + d := json.NewDecoder(strings.NewReader(raw)) + d.UseNumber() + if err := d.Decode(&doc); err != nil { + t.Fatal(err) + } + return doc +} + +func TestValidateResponse(t *testing.T) { + tests := []struct { + name string + raw string + refs []string + ok bool + }{ + {"valid", validResponse, []string{"A"}, true}, + {"no_requested_refs", validResponse, nil, false}, + {"duplicate_requested_ref", validResponse, []string{"A", "A"}, false}, + {"missing_requested_ref", validResponse, []string{"A", "B"}, false}, + {"empty_result", `{"results":{}}`, []string{"A"}, false}, + {"query_error_in_http_200", strings.Replace(validResponse, `"status":200`, `"status":200,"error":"query failed"`, 1), []string{"A"}, false}, + {"query_failed_status", strings.Replace(validResponse, `"status":200`, `"status":400`, 1), []string{"A"}, false}, + {"empty_frames", `{"results":{"A":{"status":200,"frames":[]}}}`, []string{"A"}, false}, + {"empty_rows", strings.Replace(validResponse, `[[1000,2000],[9007199254740993,2]]`, `[[],[]]`, 1), []string{"A"}, false}, + {"ragged_columns", strings.Replace(validResponse, `[9007199254740993,2]`, `[2]`, 1), []string{"A"}, false}, + {"missing_column", strings.Replace(validResponse, `[[1000,2000],[9007199254740993,2]]`, `[[1000,2000]]`, 1), []string{"A"}, false}, + {"wrong_frame_ref", strings.Replace(validResponse, `"refId":"A"`, `"refId":"B"`, 1), []string{"A"}, false}, + {"ragged_nanos", strings.Replace(validResponse, `"nanos":[[1,2],null]`, `"nanos":[[1],null]`, 1), []string{"A"}, false}, + {"wrong_nanos_width", strings.Replace(validResponse, `"nanos":[[1,2],null]`, `"nanos":[[1,2]]`, 1), []string{"A"}, false}, + {"all_null_values", strings.Replace(validResponse, `[9007199254740993,2]`, `[null,null]`, 1), []string{"A"}, false}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + err := validateResponse(decodeFixture(t, tt.raw), tt.refs) + if (err == nil) != tt.ok { + t.Fatalf("error = %v, want success %v", err, tt.ok) + } + }) + } +} + +func TestCompareResponses(t *testing.T) { + for _, tt := range []struct{ name, old, replacement string }{ + {"integer_precision", "9007199254740993", "9007199254740992"}, + {"timestamp_precision", `"nanos":[[1,2],null]`, `"nanos":[[1,3],null]`}, + {"timestamp", "[1000,2000]", "[1001,2000]"}, + {"value", "9007199254740993", "1"}, + {"label", `"job":"prom"`, `"job":"other"`}, + {"field_name", `"name":"count"`, `"name":"other"`}, + {"field_type", `"type":"number"`, `"type":"string"`}, + {"row_order", "[1000,2000]", "[2000,1000]"}, + } { + t.Run(tt.name, func(t *testing.T) { + a := decodeFixture(t, validResponse) + b := decodeFixture(t, strings.Replace(validResponse, tt.old, tt.replacement, 1)) + if compareResponses(a, b) == nil { + t.Fatal("changed data passed comparison") + } + }) + } + if err := compareResponses(decodeFixture(t, validResponse), decodeFixture(t, validResponse)); err != nil { + t.Fatal(err) + } +} + +func TestGrafanaRequest(t *testing.T) { + for _, tt := range []struct { + name, body string + status int + ok bool + }{ + {"valid", validResponse, 200, true}, + {"http_error", validResponse, 500, false}, + {"html_login", "Login", 200, false}, + {"truncated", `{"results":`, 200, false}, + {"trailing_json", validResponse + `{}`, 200, false}, + } { + t.Run(tt.name, func(t *testing.T) { + s := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost || r.URL.Path != "/api/ds/query" { + t.Errorf("unexpected request: %s %s", r.Method, r.URL.Path) + } + w.WriteHeader(tt.status) + _, _ = io.WriteString(w, tt.body) + })) + defer s.Close() + g := grafanaClient{s.URL, &http.Client{Timeout: time.Second}} + var doc queryResponse + err := g.request(http.MethodPost, "/api/ds/query", map[string]any{"queries": []any{}}, &doc) + if (err == nil) != tt.ok { + t.Fatalf("error = %v, want success %v", err, tt.ok) + } + if tt.ok && doc.Results["A"].Frames[0].Data.Values[1][0] != json.Number("9007199254740993") { + t.Fatal("lost integer precision") + } + }) + } +} diff --git a/integration/greptimedb/http_sql_test.go b/integration/greptimedb/http_sql_test.go new file mode 100644 index 000000000..04c32f5be --- /dev/null +++ b/integration/greptimedb/http_sql_test.go @@ -0,0 +1,293 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 greptimedb_test + +import ( + "bytes" + "crypto/sha256" + "encoding/json" + "fmt" + "io" + "math/big" + "net/http" + "net/url" + "os" + "path/filepath" + "reflect" + "runtime" + "strings" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/proxy/headers" +) + +type httpSQLResponse struct { + Status int `json:"status"` + Headers http.Header `json:"headers"` + Body json.RawMessage `json:"body"` + Text string `json:"text,omitempty"` +} + +func fetchHTTPSQL(endpoint, method string, values url.Values, extra http.Header) (httpSQLResponse, error) { + var out httpSQLResponse + var body io.Reader + if method == http.MethodPost { + body = strings.NewReader(values.Encode()) + } else { + endpoint += "?" + values.Encode() + } + r, err := http.NewRequest(method, endpoint, body) + if err != nil { + return out, err + } + r.Header = extra.Clone() + if r.Header == nil { + r.Header = make(http.Header) + } + r.SetBasicAuth(envOr("GREPTIMEDB_USER", "grafana_ro"), envOr("GREPTIMEDB_PASSWORD", "trickster-dev-grafana")) + if method == http.MethodPost { + r.Header.Set("Content-Type", "application/x-www-form-urlencoded") + } + resp, err := (&http.Client{Timeout: 30 * time.Second}).Do(r) + if err != nil { + return out, err + } + defer resp.Body.Close() + raw, err := io.ReadAll(io.LimitReader(resp.Body, 16<<20)) + if err != nil { + return out, err + } + out.Status, out.Headers = resp.StatusCode, resp.Header.Clone() + if json.Valid(raw) { + out.Body = raw + } else { + out.Text = string(raw) + } + return out, nil +} + +func normalizeHTTPNumbers(value any) any { + switch v := value.(type) { + case json.Number: + if r, ok := new(big.Rat).SetString(string(v)); ok { + return struct{ Number string }{r.RatString()} + } + case []any: + for i, x := range v { + v[i] = normalizeHTTPNumbers(x) + } + case map[string]any: + for k, x := range v { + v[k] = normalizeHTTPNumbers(x) + } + } + return value +} + +func compareHTTPSQL(left, right httpSQLResponse) error { + if left.Status != right.Status { + return fmt.Errorf("origin status %d, proxy status %d", left.Status, right.Status) + } + if left.Body == nil || right.Body == nil { + if !bytes.Equal(left.Body, right.Body) || left.Text != right.Text { + return fmt.Errorf("HTTP text response differs") + } + return nil + } + var l, r map[string]any + for _, entry := range []struct { + body []byte + out *map[string]any + }{{left.Body, &l}, {right.Body, &r}} { + d := json.NewDecoder(bytes.NewReader(entry.body)) + d.UseNumber() + if err := d.Decode(entry.out); err != nil { + return err + } + delete(*entry.out, "execution_time_ms") + } + if !reflect.DeepEqual(normalizeHTTPNumbers(l), normalizeHTTPNumbers(r)) { + return fmt.Errorf("typed HTTP SQL result differs: origin=%s proxy=%s", left.Body, right.Body) + } + return nil +} + +// TestHTTPSQLCacheEnvironment is read-only and requires the seeded developer +// database plus the built Trickster HTTP listener. Every result is retained. +func TestHTTPSQLCacheEnvironment(t *testing.T) { + if os.Getenv("TRICKSTER_GREPTIMEDB_HTTP_ACCEPTANCE") != "1" { + t.Skip("set TRICKSTER_GREPTIMEDB_HTTP_ACCEPTANCE=1 with running HTTP listeners") + } + out := os.Getenv("GREPTIMEDB_REPORT_DIR") + if out == "" { + t.Fatal("GREPTIMEDB_REPORT_DIR is required") + } + if err := os.MkdirAll(out, 0o700); err != nil { + t.Fatal(err) + } + file := filepath.Join(out, "http-sql-report.json") + f, err := os.OpenFile(file, os.O_CREATE|os.O_EXCL|os.O_WRONLY, 0o600) + if err != nil { + t.Fatal(err) + } + if err := f.Close(); err != nil { + t.Fatal(err) + } + r := report{StartedAt: time.Now().UTC(), GoVersion: runtime.Version(), BuildNote: os.Getenv("GREPTIMEDB_BUILD_NOTE"), Hashes: map[string]string{}, Facts: map[string]any{}} + t.Cleanup(func() { + if err := writeJSON(file, r); err != nil { + t.Error(err) + } + }) + raw, err := os.ReadFile(filepath.Join(environment, "seed-data", "seed-window.env")) + if err != nil { + t.Fatal(err) + } + seed, err := readSeed(raw) + if err != nil { + t.Fatal(err) + } + r.Hashes["seed-window.env"] = fmt.Sprintf("%x", sha256.Sum256(raw)) + r.To = time.Unix(seed["SEED_EPOCH"], 0).UTC().Truncate(24 * time.Hour) + r.From = r.To.Add(-48 * time.Hour) + origin := strings.TrimRight(envOr("GREPTIMEDB_HTTP_URL", "http://127.0.0.1:4000"), "/") + "/v1/sql" + proxy := strings.TrimRight(envOr("GREPTIMEDB_PROXY_HTTP_URL", "http://127.0.0.1:8480/greptimedb1"), "/") + "/v1/sql" + nonce := fmt.Sprint(time.Now().UnixNano()) + run := func(name, method, statement string, extra url.Values, hdr http.Header, engine, status string, reference ...string) { + t.Run(name, func(t *testing.T) { + c := check{Name: name, Status: "PASS"} + defer func() { + if t.Failed() { + c.Status = "FAIL" + } + r.Checks = append(r.Checks, c) + }() + values := url.Values{"sql": {statement}, "db": {"public"}} + for k, v := range extra { + values[k] = v + } + originValues := url.Values{} + for k, v := range values { + originValues[k] = v + } + if len(reference) != 0 { + originValues.Set("sql", reference[0]) + } + want, err := fetchHTTPSQL(origin, method, originValues, hdr) + if err != nil { + t.Fatal(err) + } + got, err := fetchHTTPSQL(proxy, method, values, hdr) + if err != nil { + t.Fatal(err) + } + if err := writeJSON(filepath.Join(out, name+".json"), map[string]any{"query": statement, "reference_query": originValues.Get("sql"), "origin": want, "proxy": got}); err != nil { + t.Fatal(err) + } + if want.Status != http.StatusOK { + t.Fatalf("valid fixture query failed at origin: %d %s %s", want.Status, want.Body, want.Text) + } + if err := compareHTTPSQL(want, got); err != nil { + c.Detail = err.Error() + t.Error(err) + } + actualEngine, actualStatus := headers.ParseResultEngineStatus(got.Headers.Get(headers.NameTricksterResult)) + if engine != "" && (actualEngine != engine || actualStatus != status) { + t.Errorf("cache path %s/%s, want %s/%s", actualEngine, actualStatus, engine, status) + } + }) + } + for _, method := range []string{http.MethodGet, http.MethodPost} { + for _, kind := range []string{"date_bin", "interval", "epoch", "empty"} { + prefix := strings.ToLower(method) + "_" + kind + statement := func(from, to time.Time) string { + bucket := "date_bin('15m', pickup_datetime)" + if kind == "interval" { + bucket = "date_bin(INTERVAL '15 minutes', pickup_datetime)" + } + predicate := fmt.Sprintf("pickup_datetime >= '%s' AND pickup_datetime < '%s'", from.Format(time.RFC3339Nano), to.Format(time.RFC3339Nano)) + if kind == "epoch" { + bucket = "floor(pickup_epoch / 900) * 900" + predicate = fmt.Sprintf("pickup_epoch >= %d AND pickup_epoch < %d", from.Unix(), to.Unix()) + } + if kind == "empty" { + predicate += " AND cab_type = 'missing_http_fixture'" + } + return fmt.Sprintf("SELECT %s AS http_%s_%s, cab_type, count(*) AS trips, count(*) + 9007199254740993 AS exact FROM trips WHERE %s GROUP BY 1,2 ORDER BY 1 DESC,2", bucket, prefix, nonce, predicate) + } + first := statement(r.From.Add(time.Hour), r.To.Add(-time.Hour)) + wide := statement(r.From, r.To) + run(prefix+"_miss", method, first, nil, nil, "DeltaProxyCache", "kmiss") + run(prefix+"_hit", method, first, nil, nil, "DeltaProxyCache", "hit") + run(prefix+"_partial", method, wide, nil, nil, "DeltaProxyCache", "phit") + run(prefix+"_wide_hit", method, wide, nil, nil, "DeltaProxyCache", "hit") + if kind == "date_bin" { + for _, tz := range []struct{ name, value string }{{"utc", "UTC"}, {"offset", "+08:00"}} { + hdr := http.Header{"X-Greptime-Timezone": {tz.value}} + run(prefix+"_"+tz.name+"_miss", method, wide, nil, hdr, "DeltaProxyCache", "kmiss") + run(prefix+"_"+tz.name+"_hit", method, wide, nil, hdr, "DeltaProxyCache", "hit") + } + for _, shape := range []struct { + name string + params url.Values + }{{"limit", url.Values{"limit": {"1"}}}, {"csv", url.Values{"format": {"csv"}}}} { + run(prefix+"_"+shape.name+"_miss", method, wide, shape.params, nil, "ObjectProxyCache", "kmiss") + status := "hit" + if method == http.MethodGet { + // GreptimeDB does not mark authenticated GET responses shareable. + status = "kmiss" + } + run(prefix+"_"+shape.name+"_repeat", method, wide, shape.params, nil, "ObjectProxyCache", status) + } + // Compare the unaligned input against explicit complete-bucket SQL, + // without using the provider analyzer to construct the oracle. + unaligned := statement(r.From.Add(time.Second), r.To.Add(-time.Second)) + aligned := statement(r.From.Add(15*time.Minute), r.To.Add(-15*time.Minute)) + run(prefix+"_unaligned", method, unaligned, nil, nil, "DeltaProxyCache", "hit", aligned) + run(prefix+"_unaligned_cold", method, unaligned, nil, http.Header{"Cache-Control": {"no-cache"}}, "DeltaProxyCache", "purge", aligned) + inclusive := strings.Replace(wide, "pickup_datetime <", "pickup_datetime <=", 1) + run(prefix+"_inclusive_upper", method, inclusive, nil, nil, "DeltaProxyCache", "kmiss", wide) + floating := strings.Replace(wide, "count(*) AS trips", "avg(total_amount) AS trips", 1) + run(prefix+"_float_miss", method, floating, nil, nil, "DeltaProxyCache", "kmiss") + run(prefix+"_float_hit", method, floating, nil, nil, "DeltaProxyCache", "hit") + nonfinite := strings.Replace(wide, "count(*) AS trips", "max(CAST('NaN' AS DOUBLE)) AS trips", 1) + nonfinite = strings.Replace(nonfinite, "ORDER BY 1 DESC,2", "ORDER BY 3 DESC NULLS LAST,1,2", 1) + run(prefix+"_nonfinite_order", method, nonfinite, nil, nil, "HTTPProxy", "proxy-only") + valueless := strings.Replace(wide, ", count(*) AS trips, count(*) + 9007199254740993 AS exact", "", 1) + run(prefix+"_valueless", method, valueless, nil, nil, "HTTPProxy", "proxy-only") + } + } + } + run("multiple", "POST", "SELECT 1; SELECT 2", nil, nil, "HTTPProxy", "proxy-only") + run("metadata", "POST", "SHOW TABLES", nil, nil, "HTTPProxy", "proxy-only") +} + +func TestCompareHTTPSQLNumbers(t *testing.T) { + makeResponse := func(value string) httpSQLResponse { + return httpSQLResponse{Status: 200, Body: json.RawMessage(`{"output":[` + value + `],"execution_time_ms":1}`)} + } + if err := compareHTTPSQL(makeResponse("1e3"), makeResponse("1000")); err != nil { + t.Fatal(err) + } + if err := compareHTTPSQL(makeResponse("9007199254740993"), makeResponse("9007199254740992")); err == nil { + t.Fatal("integer precision loss was hidden") + } + if err := compareHTTPSQL(makeResponse("7"), makeResponse(`"7"`)); err == nil { + t.Fatal("numeric and string values were conflated") + } +} diff --git a/integration/greptimedb/lifecycle_test.go b/integration/greptimedb/lifecycle_test.go new file mode 100644 index 000000000..a882eeef1 --- /dev/null +++ b/integration/greptimedb/lifecycle_test.go @@ -0,0 +1,125 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 greptimedb_test + +import ( + "bytes" + "os" + "os/exec" + "path/filepath" + "runtime" + "strings" + "testing" +) + +func TestDeveloperSeedStartupOrder(t *testing.T) { + if runtime.GOOS == "windows" { + t.Skip("developer Docker lifecycle script uses Linux paths") + } + bash, err := exec.LookPath("bash") + if err != nil { + t.Skip("bash is required for the developer lifecycle script") + } + script, err := os.ReadFile("../../hack/developer-seed-data.sh") + if err != nil { + t.Fatal(err) + } + for _, tt := range []struct { + name, target, mode string + wantFailure bool + }{ + {"running_seeders", "greptimedb", "active", false}, + {"no_running_seeders", "greptimedb", "idle", false}, + {"startup_failure", "greptimedb", "failed", true}, + {"wait_failure", "greptimedb", "wait-error", true}, + {"listing_failure", "greptimedb", "list-error", true}, + {"graphite_only", "graphite", "active", false}, + } { + t.Run(tt.name, func(t *testing.T) { + root := t.TempDir() + for _, dir := range []string{"hack", "bin", "docs/developer/environment"} { + if err := os.MkdirAll(filepath.Join(root, dir), 0o755); err != nil { + t.Fatal(err) + } + } + write := func(name string, body []byte) { + t.Helper() + if err := os.WriteFile(filepath.Join(root, name), body, 0o755); err != nil { + t.Fatal(err) + } + } + write("hack/developer-seed-data.sh", bytes.ReplaceAll(script, []byte("\r\n"), []byte("\n"))) + write("bin/docker", []byte(`#!/bin/bash +set -eu +printf '%s\n' "$*" >> "$DOCKER_LOG" +case "$*" in + 'compose ps -q --status running '*) + test "$MOCK_MODE" != list-error || exit 9 + test "$MOCK_MODE" != idle || exit 0 + printf 'startup-one\nstartup-two\n' + ;; + 'wait startup-one') printf '0\n' ;; + 'wait startup-two') + test "$MOCK_MODE" != wait-error || exit 8 + if test "$MOCK_MODE" = failed; then printf '7\n'; else printf '0\n'; fi + ;; +esac +`)) + logPath := filepath.Join(root, "docker.log") + cmd := exec.Command(bash, filepath.Join(root, "hack/developer-seed-data.sh")) + cmd.Env = append(os.Environ(), "PATH="+filepath.Join(root, "bin")+string(os.PathListSeparator)+os.Getenv("PATH"), + "DOCKER_LOG="+logPath, "MOCK_MODE="+tt.mode, "SEED_TARGET="+tt.target) + output, err := cmd.CombinedOutput() + if (err != nil) != tt.wantFailure { + t.Fatalf("error = %v, want failure %v; output: %s", err, tt.wantFailure, output) + } + data, err := os.ReadFile(logPath) + if err != nil { + t.Fatal(err) + } + log := string(data) + lines := strings.Split(strings.TrimSpace(log), "\n") + if !strings.HasPrefix(lines[0], "compose ps -q --status running ") { + t.Fatalf("first command must inspect startup seeders: %s", log) + } + if tt.target == "graphite" { + if lines[0] != "compose ps -q --status running graphite_seed" || strings.Contains(log, "run --rm seed_data_generate") { + t.Fatalf("graphite-only reload touched the trips fixture: %s", log) + } + } else { + for _, service := range []string{"seed_data_generate", "clickhouse_seed", "mysql_seed", "timescaledb_seed", "greptimedb_seed", "druid_seed"} { + if !strings.Contains(lines[0], service) { + t.Fatalf("shared fixture consumer %s was not checked: %s", service, log) + } + } + } + if tt.wantFailure { + if strings.Contains(log, "compose up") || strings.Contains(log, "compose run") || strings.Contains(log, "compose stop") { + t.Fatalf("startup failure must prevent mutations: %s", log) + } + return + } + if tt.mode == "idle" { + if strings.Contains(log, "wait startup-") { + t.Fatalf("wait called without running seeders: %s", log) + } + } else if len(lines) < 4 || lines[1] != "wait startup-one" || lines[2] != "wait startup-two" { + t.Fatalf("mutations preceded startup completion: %s", log) + } + }) + } +} diff --git a/integration/greptimedb/protocol_test.go b/integration/greptimedb/protocol_test.go new file mode 100644 index 000000000..c50b8fb4d --- /dev/null +++ b/integration/greptimedb/protocol_test.go @@ -0,0 +1,311 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 greptimedb_test + +import ( + "context" + "database/sql" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "net/url" + "strings" + "testing" + "time" + + "github.com/go-sql-driver/mysql" + "github.com/jackc/pgx/v5/pgconn" +) + +const pgTypeQuery = `SELECT CAST('2026-09-20 03:04:05.123456789' AS TIMESTAMP(0)) AS ts0, +CAST('2026-09-20 03:04:05.123456789' AS TIMESTAMP(3)) AS ts3, +CAST('2026-09-20 03:04:05.123456789' AS TIMESTAMP(6)) AS ts6, +CAST('2026-09-20 03:04:05.123456789' AS TIMESTAMP(9)) AS ts9, +CAST('2026-09-20' AS DATE) AS d, CAST(1 AS SMALLINT) AS i2, CAST(2 AS INT) AS i4, +CAST(9007199254740993 AS BIGINT) AS i8, CAST(1.25 AS FLOAT) AS f4, CAST(2.5 AS DOUBLE) AS f8` + +func emptyResult(result *pgconn.Result) bool { + return result != nil && result.Err == nil && len(result.Rows) == 0 && + len(result.FieldDescriptions) == 0 && result.CommandTag.String() == "" +} + +func protocolChecks(run func(string, func() error), r *report) { + user := envOr("GREPTIMEDB_USER", "grafana_ro") + password := envOr("GREPTIMEDB_PASSWORD", "trickster-dev-grafana") + database := envOr("GREPTIMEDB_DATABASE", "public") + u := url.URL{ + Scheme: "postgres", User: url.UserPassword(user, password), + Host: envOr("GREPTIMEDB_PG_ADDR", "127.0.0.1:4003"), Path: "/" + database, + RawQuery: "sslmode=disable&connect_timeout=10", + } + withPG := func(fn func(context.Context, *pgconn.PgConn) error) error { + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + c, err := pgconn.Connect(ctx, u.String()) + if err != nil { + return err + } + defer c.Close(ctx) + return fn(ctx, c) + } + withMySQL := func(fn func(context.Context, *sql.DB) error) error { + cfg := mysql.NewConfig() + cfg.User, cfg.Passwd, cfg.Net, cfg.DBName = user, password, "tcp", database + cfg.Addr = envOr("GREPTIMEDB_MYSQL_ADDR", "127.0.0.1:4002") + cfg.Timeout, cfg.ReadTimeout, cfg.WriteTimeout = 10*time.Second, 30*time.Second, 15*time.Second + db, err := sql.Open("mysql", cfg.FormatDSN()) + if err != nil { + return err + } + defer db.Close() + db.SetMaxOpenConns(1) + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + return fn(ctx, db) + } + run("pg_startup_and_empty_queries", func() error { + return withPG(func(ctx context.Context, c *pgconn.PgConn) error { + status := map[string]string{} + for _, key := range []string{"server_version", "server_encoding", "client_encoding", "DateStyle", "integer_datetimes", "TimeZone"} { + status[key] = c.ParameterStatus(key) + } + r.Facts["pgconn_parameter_status_not_grafana_capture"] = status + for _, query := range []string{"", "-- ping", "/* comment */", ";", "-- comment\n;"} { + results, err := c.Exec(ctx, query).ReadAll() + if err != nil { + return fmt.Errorf("simple empty query %q: %w", query, err) + } + // pgconn exposes EmptyQueryResponse as one empty Result. + if len(results) != 1 || !emptyResult(results[0]) || c.TxStatus() != 'I' { + return fmt.Errorf("simple empty query %q: unexpected results/status", query) + } + } + return c.Ping(ctx) + }) + }) + run("pg_extended_query", func() error { + return withPG(func(ctx context.Context, c *pgconn.PgConn) error { + result := c.ExecParams(ctx, "SELECT $1::BIGINT AS answer", [][]byte{[]byte("42")}, []uint32{20}, []int16{0}, []int16{0}).Read() + if result.Err != nil { + return result.Err + } + if len(result.Rows) != 1 || len(result.Rows[0]) != 1 || string(result.Rows[0][0]) != "42" || + len(result.FieldDescriptions) != 1 || result.FieldDescriptions[0].DataTypeOID != 20 || c.TxStatus() != 'I' { + return fmt.Errorf("unexpected extended query result") + } + result = c.ExecParams(ctx, "-- ping", nil, nil, nil, nil).Read() + if !emptyResult(result) || c.TxStatus() != 'I' { + return fmt.Errorf("extended comment-only query: %v", result.Err) + } + return nil + }) + }) + run("pg_bucket_sql", func() error { + return withPG(func(ctx context.Context, c *pgconn.PgConn) error { + return checkBucketSQL(pgTextQuery(ctx, c), r, "pg") + }) + }) + run("pg_quoting", func() error { + return withPG(func(ctx context.Context, c *pgconn.PgConn) error { + return checkQuoting(pgTextQuery(ctx, c), r, "pg") + }) + }) + run("pg_cancel_capability", func() error { + return withPG(func(_ context.Context, c *pgconn.PgConn) error { + zeroKey := len(c.SecretKey()) == 4 + for _, b := range c.SecretKey() { + zeroKey = zeroKey && b == 0 + } + r.Facts["pg_cancel_capability"] = map[string]any{ + "pid_zero": c.PID() == 0, "key_zero": zeroKey, + "supports_cancel": false, "basis": "upstream uses zero cancellation credentials and the default no-op cancel handler", + } + if c.PID() != 0 || !zeroKey { + return fmt.Errorf("upstream cancellation credentials changed; re-evaluate support before setting conformance flags") + } + return nil + }) + }) + run("pg_transaction_compatibility", func() error { + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + cfg, err := pgconn.ParseConfig(u.String()) + if err != nil { + return err + } + var notices []string + cfg.OnNotice = func(_ *pgconn.PgConn, notice *pgconn.Notice) { notices = append(notices, notice.Message) } + c, err := pgconn.ConnectConfig(ctx, cfg) + if err != nil { + return err + } + defer c.Close(ctx) + var trace []map[string]any + defer func() { + r.Facts["pg_transaction_compatibility"] = map[string]any{"steps": trace, "notices": notices, "supports_transactions": false} + }() + for _, step := range []struct { + sql, value string + status byte + wantErr bool + }{ + {"BEGIN", "", 'T', false}, + {"SELECT 1", "1", 'T', false}, + {"SELECT trickster_acceptance_missing_column FROM trips", "", 'E', true}, + {"ROLLBACK", "", 'I', false}, + {"SELECT 42", "42", 'I', false}, + } { + results, err := c.Exec(ctx, step.sql).ReadAll() + item := map[string]any{"sql": step.sql, "status": string(c.TxStatus())} + var pgErr *pgconn.PgError + if errors.As(err, &pgErr) { + item["sqlstate"] = pgErr.Code + } + trace = append(trace, item) + if (err != nil) != step.wantErr || (step.wantErr && pgErr == nil) { + return fmt.Errorf("%s: unexpected error %v", step.sql, err) + } + if c.TxStatus() != step.status { + return fmt.Errorf("%s: status %c, want %c", step.sql, c.TxStatus(), step.status) + } + if step.value != "" { + if len(results) != 1 || len(results[0].Rows) != 1 || len(results[0].Rows[0]) != 1 || string(results[0].Rows[0][0]) != step.value { + return fmt.Errorf("%s: unexpected recovery query result", step.sql) + } + item["value"] = step.value + } + } + for _, notice := range notices { + if strings.Contains(notice, "transaction is not supported") && c.TxStatus() == 'I' { + return nil + } + } + return fmt.Errorf("missing upstream no-transactions warning or rollback did not restore idle status") + }) + run("pg_type_oids_and_text", func() error { + return withPG(func(ctx context.Context, c *pgconn.PgConn) error { + results, err := c.Exec(ctx, pgTypeQuery).ReadAll() + if err != nil { + return err + } + wantOID := []uint32{1114, 1114, 1114, 1114, 1082, 21, 23, 20, 700, 701} + // Pin the observed upstream wire rendering, not Arrow's stored precision. + // The PG encoder renders six fractional digits even for TIMESTAMP(9). + wantText := []string{"2026-09-20 03:04:05.000000", "2026-09-20 03:04:05.123000", "2026-09-20 03:04:05.123456", "2026-09-20 03:04:05.123456", "2026-09-20", "1", "2", "9007199254740993", "1.25", "2.5"} + if len(results) != 1 || len(results[0].Rows) != 1 || len(results[0].FieldDescriptions) != len(wantOID) || len(results[0].Rows[0]) != len(wantOID) { + return fmt.Errorf("unexpected typed query result shape") + } + var observed []map[string]any + for i, f := range results[0].FieldDescriptions { + observed = append(observed, map[string]any{"name": f.Name, "oid": f.DataTypeOID, "format": f.Format, "text": string(results[0].Rows[0][i])}) + } + r.Facts["pgconn_type_metadata_not_grafana_capture"] = observed + for i, f := range results[0].FieldDescriptions { + if f.DataTypeOID != wantOID[i] || f.Format != 0 || string(results[0].Rows[0][i]) != wantText[i] { + return fmt.Errorf("field %s: oid=%d format=%d text=%q, want oid=%d text=%q", f.Name, f.DataTypeOID, f.Format, results[0].Rows[0][i], wantOID[i], wantText[i]) + } + } + return nil + }) + }) + run("http_sql", func() error { + form := url.Values{"db": {database}, "sql": {"SELECT CAST(9007199254740993 AS BIGINT) AS answer"}} + endpoint := strings.TrimRight(envOr("GREPTIMEDB_HTTP_URL", "http://127.0.0.1:4000"), "/") + "/v1/sql" + req, err := http.NewRequest(http.MethodPost, endpoint, strings.NewReader(form.Encode())) + if err != nil { + return err + } + req.SetBasicAuth(user, password) + req.Header.Set("Content-Type", "application/x-www-form-urlencoded") + resp, err := (&http.Client{Timeout: 15 * time.Second}).Do(req) + if err != nil { + return err + } + defer resp.Body.Close() + var body struct { + Code int `json:"code"` + Error string `json:"error"` + Output []struct { + Records struct { + Rows [][]json.Number `json:"rows"` + } `json:"records"` + } `json:"output"` + } + if err := json.NewDecoder(io.LimitReader(resp.Body, 1<<20)).Decode(&body); err != nil { + return err + } + if resp.StatusCode != 200 || body.Code != 0 || body.Error != "" || len(body.Output) != 1 || + len(body.Output[0].Records.Rows) != 1 || len(body.Output[0].Records.Rows[0]) != 1 || + body.Output[0].Records.Rows[0][0] != json.Number("9007199254740993") { + return fmt.Errorf("unexpected HTTP SQL response (HTTP %d): %+v", resp.StatusCode, body) + } + return nil + }) + run("mysql_ping_and_query", func() error { + return withMySQL(func(ctx context.Context, db *sql.DB) error { + if err := db.PingContext(ctx); err != nil { + return err + } + var version string + if err := db.QueryRowContext(ctx, "SELECT version()").Scan(&version); err != nil { + return err + } + r.Facts["mysql_version"] = version + var answer int64 + if err := db.QueryRowContext(ctx, "SELECT CAST(9007199254740993 AS BIGINT)").Scan(&answer); err != nil { + return err + } + if answer != 9007199254740993 { + return fmt.Errorf("unexpected MySQL value %d", answer) + } + return nil + }) + }) + run("mysql_bucket_sql", func() error { + return withMySQL(func(ctx context.Context, db *sql.DB) error { + return checkBucketSQL(mysqlTextQuery(ctx, db), r, "mysql") + }) + }) + run("mysql_quoting", func() error { + return withMySQL(func(ctx context.Context, db *sql.DB) error { + return checkQuoting(mysqlTextQuery(ctx, db), r, "mysql") + }) + }) +} + +func TestEmptyResult(t *testing.T) { + for _, tt := range []struct { + name string + result *pgconn.Result + want bool + }{ + {"empty_query_response", &pgconn.Result{}, true}, + {"missing_response", nil, false}, + {"query_error", &pgconn.Result{Err: fmt.Errorf("query failed")}, false}, + {"row", &pgconn.Result{Rows: [][][]byte{{[]byte("1")}}}, false}, + {"schema", &pgconn.Result{FieldDescriptions: []pgconn.FieldDescription{{Name: "a"}}}, false}, + {"command_complete", &pgconn.Result{CommandTag: pgconn.NewCommandTag("SELECT 0")}, false}, + } { + t.Run(tt.name, func(t *testing.T) { + if got := emptyResult(tt.result); got != tt.want { + t.Fatalf("emptyResult = %v, want %v", got, tt.want) + } + }) + } +} diff --git a/integration/greptimedb/proxy_test.go b/integration/greptimedb/proxy_test.go new file mode 100644 index 000000000..5319a906e --- /dev/null +++ b/integration/greptimedb/proxy_test.go @@ -0,0 +1,482 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 greptimedb_test + +import ( + "context" + "crypto/sha256" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "net/http/httptest" + "net/url" + "os" + "path/filepath" + "reflect" + "runtime" + "strconv" + "strings" + "testing" + "time" + + "github.com/jackc/pgx/v5/pgconn" + "github.com/prometheus/common/expfmt" + "github.com/prometheus/common/model" +) + +// TestProxyEnvironment compares the same origin through different transports. +// Cross-engine rounding allowances from the direct suite never apply here. +func TestProxyEnvironment(t *testing.T) { + if os.Getenv("TRICKSTER_GREPTIMEDB_PROXY_ACCEPTANCE") != "1" { + t.Skip("set TRICKSTER_GREPTIMEDB_PROXY_ACCEPTANCE=1 with running proxy listeners") + } + out := os.Getenv("GREPTIMEDB_REPORT_DIR") + if out == "" { + t.Fatal("GREPTIMEDB_REPORT_DIR is required") + } + if err := os.MkdirAll(out, 0o700); err != nil { + t.Fatal(err) + } + f, err := os.OpenFile(filepath.Join(out, "proxy-report.json"), os.O_CREATE|os.O_EXCL|os.O_WRONLY, 0o600) + if err != nil { + t.Fatal(err) + } + if err = f.Close(); err != nil { + t.Fatal(err) + } + r := report{ + StartedAt: time.Now().UTC(), GoVersion: runtime.Version(), BuildNote: os.Getenv("GREPTIMEDB_BUILD_NOTE"), + Hashes: map[string]string{}, Facts: map[string]any{}, + } + t.Cleanup(func() { + if err := writeJSON(filepath.Join(out, "proxy-report.json"), r); err != nil { + t.Error(err) + } + }) + run := func(name string, fn func() error) { + t.Run(name, func(t *testing.T) { + c := check{Name: name, Status: "PASS"} + if err := fn(); err != nil { + c.Status, c.Detail = "FAIL", err.Error() + t.Error(err) + } + r.Checks = append(r.Checks, c) + }) + } + raw, err := os.ReadFile(filepath.Join(environment, "seed-data", "seed-window.env")) + if err != nil { + t.Fatal(err) + } + seed, err := readSeed(raw) + if err != nil { + t.Fatal(err) + } + r.Hashes["seed-window.env"] = fmt.Sprintf("%x", sha256.Sum256(raw)) + r.To = time.Unix(seed["SEED_EPOCH"], 0).UTC().Truncate(24 * time.Hour) + r.From = r.To.Add(-48 * time.Hour) + raw, err = os.ReadFile(filepath.Join(environment, "dashboards", "trickster-greptimedb.json")) + if err != nil { + t.Fatal(err) + } + var db dashboard + if err = json.Unmarshal(raw, &db); err != nil { + t.Fatal(err) + } + r.Hashes["greptimedb-dashboard"] = fmt.Sprintf("%x", sha256.Sum256(raw)) + g := grafanaClient{strings.TrimRight(envOr("GRAFANA_URL", "http://127.0.0.1:3000"), "/"), &http.Client{Timeout: 30 * time.Second}} + user, password := envOr("GREPTIMEDB_USER", "grafana_ro"), envOr("GREPTIMEDB_PASSWORD", "trickster-dev-grafana") + for _, mode := range []struct{ name, addr, uid, backend string }{ + {"terminated", envOr("GREPTIMEDB_PROXY_PG_ADDR", "127.0.0.1:8489"), envOr("GREPTIMEDB_PROXY_SQL_UID", "ds_greptimedb_trickster"), envOr("GREPTIMEDB_PROXY_BACKEND", "greptimedb1")}, + {"passthrough", os.Getenv("GREPTIMEDB_PASSTHROUGH_PG_ADDR"), os.Getenv("GREPTIMEDB_PASSTHROUGH_SQL_UID"), envOr("GREPTIMEDB_PASSTHROUGH_BACKEND", "greptimedb-pass")}, + } { + if mode.addr == "" || mode.uid == "" { + r.Checks = append(r.Checks, check{mode.name, "UNVERIFIED", "No passthrough listener/datasource supplied"}) + continue + } + run(mode.name+"_pgwire", func() error { return comparePGProxy(mode.addr, user, password) }) + run(mode.name+"_grafana_health", func() error { + var health struct{ Status, Message string } + if err := g.request(http.MethodGet, "/api/datasources/uid/"+url.PathEscape(mode.uid)+"/health", nil, &health); err != nil { + return err + } + if health.Status != "OK" { + return fmt.Errorf("datasource status %q: %s", health.Status, health.Message) + } + return nil + }) + for _, panel := range db.Panels { + if !((panel.ID >= 1 && panel.ID <= 6) || (panel.ID >= 20 && panel.ID <= 25)) { + continue + } + cacheMode := "delta" + if panel.ID == 5 || panel.ID == 22 || panel.ID == 23 || panel.ID == 24 { + cacheMode = "object" + } + for scenario, shift := range []time.Duration{0, 0, time.Hour} { + run(fmt.Sprintf("%s_panel_%d_%d", mode.name, panel.ID, scenario), func() error { + before, err := sqlCacheCounts(g.client, mode.backend, cacheMode) + if err != nil { + return err + } + var docs []queryResponse + for _, uid := range []string{envOr("GREPTIMEDB_SQL_UID", "ds_greptimedb_direct"), mode.uid} { + doc, err := g.query(panel.Targets, uid, "grafana-postgresql-datasource", r.From.Add(shift), r.To.Add(shift), 5*time.Minute) + if writeErr := writeJSON(filepath.Join(out, fmt.Sprintf("proxy-%s-panel-%d-%d-%s.json", mode.name, panel.ID, scenario, uid)), doc); writeErr != nil { + return writeErr + } + if err != nil { + return err + } + docs = append(docs, doc) + } + after, err := sqlCacheCounts(g.client, mode.backend, cacheMode) + if err != nil { + return err + } + status, err := cacheTransition(before, after, cacheMode, scenario) + r.Facts[fmt.Sprintf("%s_panel_%d_%d_cache", mode.name, panel.ID, scenario)] = map[string]any{"mode": cacheMode, "status": status, "before": before, "after": after} + return errors.Join(err, compareResponses(docs[0], docs[1])) + }) + } + } + } + originHTTP := strings.TrimRight(envOr("GREPTIMEDB_HTTP_URL", "http://127.0.0.1:4000"), "/") + proxyHTTP := strings.TrimRight(envOr("GREPTIMEDB_PROXY_HTTP_URL", "http://127.0.0.1:8480/greptimedb1"), "/") + to := time.Now().UTC().Add(-30 * time.Second).Truncate(15 * time.Second) + from := to.Add(-time.Minute) + r.Facts["promql_window"] = map[string]any{"from": from, "to": to, "step_seconds": 15} + for _, query := range []struct { + name, path string + form url.Values + }{ + {"http_sql", "/v1/sql", url.Values{"db": {"public"}, "sql": {"SELECT COUNT(*) AS rows, MIN(pickup_epoch) AS first, MAX(pickup_epoch) AS last FROM trips"}}}, + {"http_promql", "/v1/prometheus/api/v1/query_range", url.Values{"query": {`up{job="prometheus"}`}, "start": {strconv.FormatInt(from.Unix(), 10)}, "end": {strconv.FormatInt(to.Unix(), 10)}, "step": {"15"}}}, + } { + run(query.name, func() error { + want, err := proxyHTTPPayload(originHTTP+query.path, query.form, user, password, query.name == "http_sql") + if err != nil { + return err + } + if err := writeJSON(filepath.Join(out, "origin-"+query.name+".json"), want); err != nil { + return err + } + for n := range 2 { + got, err := proxyHTTPPayload(proxyHTTP+query.path, query.form, user, password, query.name == "http_sql") + if err != nil { + return err + } + if err := writeJSON(filepath.Join(out, fmt.Sprintf("proxy-%s-%d.json", query.name, n)), got); err != nil { + return err + } + if query.name == "http_promql" { + // JSON timestamps may be spelled as 1 or 1.0. Compare their + // exact numeric value without rounding large integers. + want, got = normalizeHTTPNumbers(want), normalizeHTTPNumbers(got) + } + if !reflect.DeepEqual(want, got) { + return fmt.Errorf("direct and proxied %s payloads differ", query.name) + } + } + return nil + }) + } + run("grafana_promql", func() error { + targets := []map[string]any{{"refId": "A", "expr": `up{job="prometheus"}`, "range": true, "instant": false, "interval": "15s"}} + want, err := g.query(targets, envOr("GREPTIMEDB_PROM_UID", "ds_greptimedb_prom_direct"), "prometheus", from, to, 15*time.Second) + if err != nil { + return err + } + got, err := g.query(targets, envOr("GREPTIMEDB_PROXY_PROM_UID", "ds_greptimedb_prom_trickster"), "prometheus", from, to, 15*time.Second) + if err != nil { + return err + } + return compareResponses(want, got) + }) + run("provider_metrics", func() error { + endpoint := envOr("GREPTIMEDB_PROXY_METRICS_URL", "http://127.0.0.1:8481/metrics") + resp, err := g.client.Get(endpoint) + if err != nil { + return err + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusOK { + return fmt.Errorf("metrics HTTP %d", resp.StatusCode) + } + parser := expfmt.NewTextParser(model.UTF8Validation) + families, err := parser.TextToMetricFamilies(io.LimitReader(resp.Body, 16<<20)) + if err != nil { + return err + } + var httpRequests, pgRequests float64 + for _, m := range families["trickster_proxy_requests_total"].GetMetric() { + labels := make(map[string]string) + for _, l := range m.GetLabel() { + labels[l.GetName()] = l.GetValue() + } + if labels["provider"] != "greptimedb" || labels["backend_name"] != envOr("GREPTIMEDB_PROXY_BACKEND", "greptimedb1") { + continue + } + if labels["path"] == "query" { + pgRequests += m.GetCounter().GetValue() + } + } + // The standard ReverseProxy lane records HTTP traffic at the frontend, + // whereas native pgwire requests use proxy counters. + for _, m := range families["trickster_frontend_requests_total"].GetMetric() { + labels := make(map[string]string) + for _, l := range m.GetLabel() { + labels[l.GetName()] = l.GetValue() + } + if labels["provider"] == "greptimedb" && labels["backend_name"] == envOr("GREPTIMEDB_PROXY_BACKEND", "greptimedb1") && labels["http_status"] == "2xx" { + httpRequests += m.GetCounter().GetValue() + } + } + if httpRequests == 0 || pgRequests == 0 { + return fmt.Errorf("missing HTTP or pgwire provider metrics: HTTP=%v pgwire=%v", httpRequests, pgRequests) + } + r.Facts["provider_requests"] = map[string]float64{"http": httpRequests, "pgwire": pgRequests} + return nil + }) +} + +func sqlCacheCounts(client *http.Client, backend, mode string) (map[string]float64, error) { + resp, err := client.Get(envOr("GREPTIMEDB_PROXY_METRICS_URL", "http://127.0.0.1:8481/metrics")) + if err != nil { + return nil, err + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusOK { + return nil, fmt.Errorf("metrics HTTP %d", resp.StatusCode) + } + parser := expfmt.NewTextParser(model.UTF8Validation) + families, err := parser.TextToMetricFamilies(io.LimitReader(resp.Body, 16<<20)) + if err != nil { + return nil, err + } + counts := map[string]float64{} + for _, metric := range families["trickster_sql_query_cache_total"].GetMetric() { + labels := map[string]string{} + for _, label := range metric.GetLabel() { + labels[label.GetName()] = label.GetValue() + } + if labels["backend_name"] == backend && labels["cache_mode"] == mode { + counts[labels["cache_status"]] += metric.GetCounter().GetValue() + } + } + for _, metric := range families["trickster_sql_query_rewrite_failures_total"].GetMetric() { + for _, label := range metric.GetLabel() { + if label.GetName() == "backend_name" && label.GetValue() == backend && metric.GetCounter().GetValue() > 0 { + return nil, fmt.Errorf("backend %s has SQL rewrite failures", backend) + } + } + } + return counts, nil +} + +func cacheTransition(before, after map[string]float64, mode string, scenario int) (string, error) { + status := "" + for name, value := range after { + delta := value - before[name] + if delta == 0 { + continue + } + if delta != 1 || status != "" { + return "", fmt.Errorf("expected one cache request, before=%v after=%v", before, after) + } + status = name + } + if status != "kmiss" && status != "rmiss" && status != "hit" && status != "phit" { + return status, fmt.Errorf("query did not take the %s cache path: %s", mode, status) + } + if scenario == 1 && status != "hit" { + return status, fmt.Errorf("repeat was %s, not a hit", status) + } + if scenario == 2 && mode == "delta" && status != "phit" && status != "hit" { + return status, fmt.Errorf("overlapping range was %s, not a partial hit or a warmed hit", status) + } + if scenario == 2 && mode == "object" && status != "kmiss" && status != "hit" { + return status, fmt.Errorf("object query unexpectedly reported %s", status) + } + return status, nil +} + +func TestCacheTransition(t *testing.T) { + for _, tc := range []struct { + mode, status string + phase int + pass bool + }{ + {"delta", "kmiss", 0, true}, + {"delta", "hit", 1, true}, + {"delta", "phit", 2, true}, + {"delta", "kmiss", 1, false}, + {"delta", "kmiss", 2, false}, + {"delta", "", 0, false}, + {"object", "kmiss", 2, true}, + {"object", "phit", 2, false}, + } { + _, err := cacheTransition(nil, map[string]float64{tc.status: 1}, tc.mode, tc.phase) + if (err == nil) != tc.pass { + t.Fatalf("%+v: %v", tc, err) + } + } + for _, counts := range []map[string]float64{{"hit": 2}, {"hit": 1, "kmiss": 1}, {"hit": -1}} { + if _, err := cacheTransition(nil, counts, "delta", 0); err == nil { + t.Fatalf("accepted %v", counts) + } + } +} + +func comparePGProxy(address, user, password string) error { + ctx, cancel := context.WithTimeout(context.Background(), time.Minute) + defer cancel() + connect := func(addr, pass string) (*pgconn.PgConn, error) { + u := url.URL{Scheme: "postgres", Host: addr, Path: "/" + envOr("GREPTIMEDB_DATABASE", "public"), User: url.UserPassword(user, pass), RawQuery: "sslmode=disable&connect_timeout=10"} + return pgconn.Connect(ctx, u.String()) + } + direct, err := connect(envOr("GREPTIMEDB_PG_ADDR", "127.0.0.1:4003"), password) + if err != nil { + return err + } + defer direct.Close(ctx) + proxy, err := connect(address, password) + if err != nil { + return err + } + defer proxy.Close(ctx) + bad, err := connect(address, password+"-wrong") + if err == nil { + _ = bad.Close(ctx) + return fmt.Errorf("proxy accepted a wrong password") + } + for _, sql := range []string{"", "-- ping", pgTypeQuery, "SELECT 1 AS a; SELECT 2 AS b", "SELECT COUNT(*) FROM trips"} { + want, err := direct.Exec(ctx, sql).ReadAll() + if err != nil { + return err + } + got, err := proxy.Exec(ctx, sql).ReadAll() + if err != nil { + return err + } + if !reflect.DeepEqual(want, got) { + return fmt.Errorf("pgwire results differ for %q", sql) + } + } + want := direct.ExecParams(ctx, "SELECT $1::BIGINT AS answer", [][]byte{[]byte("42")}, []uint32{20}, nil, nil).Read() + got := proxy.ExecParams(ctx, "SELECT $1::BIGINT AS answer", [][]byte{[]byte("42")}, []uint32{20}, nil, nil).Read() + if want.Err != nil || got.Err != nil || !reflect.DeepEqual(want, got) { + return fmt.Errorf("extended-protocol responses differ: direct=%v proxy=%v", want.Err, got.Err) + } + _, wantErr := direct.Exec(ctx, "SELECT __missing_proxy_column FROM trips").ReadAll() + _, gotErr := proxy.Exec(ctx, "SELECT __missing_proxy_column FROM trips").ReadAll() + var wantPG, gotPG *pgconn.PgError + if !errors.As(wantErr, &wantPG) || !errors.As(gotErr, &gotPG) || wantPG.Code != gotPG.Code { + return fmt.Errorf("error SQLSTATE differs") + } + return proxy.Ping(ctx) +} + +func proxyHTTPPayload(endpoint string, form url.Values, user, password string, sql bool) (any, error) { + req, err := http.NewRequest(http.MethodPost, endpoint, strings.NewReader(form.Encode())) + if err != nil { + return nil, err + } + req.SetBasicAuth(user, password) + req.Header.Set("Content-Type", "application/x-www-form-urlencoded") + client := &http.Client{Timeout: 30 * time.Second} + resp, err := client.Do(req) + if err != nil { + return nil, err + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusOK { + return nil, fmt.Errorf("HTTP %d from %s", resp.StatusCode, req.URL.Path) + } + d := json.NewDecoder(io.LimitReader(resp.Body, 1<<20)) + d.UseNumber() + var body map[string]any + if err := d.Decode(&body); err != nil { + return nil, err + } + if err := d.Decode(new(any)); err != io.EOF { + return nil, fmt.Errorf("trailing HTTP response data") + } + if sql { + output, ok := body["output"].([]any) + if !ok || len(output) == 0 || body["error"] != nil { + return nil, fmt.Errorf("unsuccessful SQL output") + } + for _, item := range output { + result, ok := item.(map[string]any) + if !ok { + return nil, fmt.Errorf("invalid SQL result") + } + records, ok := result["records"].(map[string]any) + if !ok { + return nil, fmt.Errorf("missing SQL records") + } + rows, ok := records["rows"].([]any) + if !ok || len(rows) == 0 { + return nil, fmt.Errorf("empty SQL records") + } + } + // Execution duration is not query data and varies between requests. + return output, nil + } + data, ok := body["data"].(map[string]any) + if body["status"] != "success" || !ok { + return nil, fmt.Errorf("unsuccessful PromQL output") + } + result, ok := data["result"].([]any) + if !ok || len(result) == 0 { + return nil, fmt.Errorf("empty PromQL output") + } + return data, nil +} + +func TestProxyHTTPPayloadValidation(t *testing.T) { + for _, tt := range []struct { + name, body string + sql, valid bool + }{ + {"SQL data", `{"output":[{"records":{"rows":[[9007199254740993]]}}],"execution_time_ms":3}`, true, true}, + {"SQL error", `{"error":"bad query","code":1000}`, true, false}, + {"empty SQL", `{"output":[]}`, true, false}, + {"empty records", `{"output":[{"records":{"rows":[]}}]}`, true, false}, + {"invalid records", `{"output":[{}]}`, true, false}, + {"PromQL data", `{"status":"success","data":{"resultType":"matrix","result":[{"values":[[1,"1"]]}]}}`, false, true}, + {"empty PromQL", `{"status":"success","data":{"result":[]}}`, false, false}, + {"PromQL error", `{"status":"error","error":"bad query"}`, false, false}, + {"trailing JSON", `{"output":[]} {}`, true, false}, + } { + t.Run(tt.name, func(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + t.Error("expected POST") + } + _, _ = io.WriteString(w, tt.body) + })) + defer server.Close() + _, err := proxyHTTPPayload(server.URL, url.Values{"sql": {"SELECT 1"}}, "user", "password", tt.sql) + if (err == nil) != tt.valid { + t.Fatalf("valid=%t, err=%v", tt.valid, err) + } + }) + } +} diff --git a/integration/greptimedb_mysql_test.go b/integration/greptimedb_mysql_test.go new file mode 100644 index 000000000..af7cda820 --- /dev/null +++ b/integration/greptimedb_mysql_test.go @@ -0,0 +1,301 @@ +/* + * Copyright 2026 The Trickster Authors + * 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 integration + +import ( + "context" + "encoding/json" + "fmt" + "net" + "net/http" + "net/url" + "os" + "strconv" + "strings" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/integration/internal/portutil" + "github.com/trickstercache/trickster/v2/pkg/backends/greptimedb" + backend "github.com/trickstercache/trickster/v2/pkg/backends/mysql" + mo "github.com/trickstercache/trickster/v2/pkg/backends/mysql/options" + bo "github.com/trickstercache/trickster/v2/pkg/backends/options" + cachemanager "github.com/trickstercache/trickster/v2/pkg/cache/manager" + cachememory "github.com/trickstercache/trickster/v2/pkg/cache/memory" + cacheoptions "github.com/trickstercache/trickster/v2/pkg/cache/options" + tkconfig "github.com/trickstercache/trickster/v2/pkg/config" + "github.com/trickstercache/trickster/v2/pkg/config/listener" + "github.com/trickstercache/trickster/v2/pkg/observability/metrics" + pgo "github.com/trickstercache/trickster/v2/pkg/proxy/pgwire/options" + + dto "github.com/prometheus/client_model/go" + "github.com/stretchr/testify/require" + vtmysql "vitess.io/vitess/go/mysql" + "vitess.io/vitess/go/sqltypes" +) + +func TestGreptimeMySQLSharedBackend(t *testing.T) { + if os.Getenv("TRICKSTER_GREPTIMEDB_MYSQL_ACCEPTANCE") != "1" { + t.Skip("set TRICKSTER_GREPTIMEDB_MYSQL_ACCEPTANCE=1 with isolated developer GreptimeDB") + } + ports, release := portutil.Reserve(t, 2) + mysqlOrigin := os.Getenv("GREPTIMEDB_MYSQL_ADDR") + if mysqlOrigin == "" { + mysqlOrigin = "127.0.0.1:4002" + } + pgOrigin := os.Getenv("GREPTIMEDB_PG_ADDR") + if pgOrigin == "" { + pgOrigin = "127.0.0.1:4003" + } + httpOrigin := os.Getenv("GREPTIMEDB_HTTP_URL") + if httpOrigin == "" { + httpOrigin = "http://127.0.0.1:4000" + } + harness := configHarness(t, func(c *tkconfig.Config) { + c.Listeners["greptime-mysql"] = &listener.Options{Protocol: listener.ProtocolMySQL, ListenAddress: "127.0.0.1", ListenPort: ports[0]} + c.Listeners["greptime-pg"] = &listener.Options{Protocol: listener.ProtocolPostgres, ListenAddress: "127.0.0.1", ListenPort: ports[1]} + o := bo.New() + o.Provider, o.OriginURL = "greptimedb", httpOrigin + o.ListenerNames = []string{"default", "greptime-mysql", "greptime-pg"} + o.AuthenticatorName, o.CacheName = "greptimedb-grafana", "mem1" + u := url.URL{Scheme: "mysql", Host: mysqlOrigin, User: url.UserPassword("grafana_ro", "trickster-dev-grafana"), Path: "/public"} + o.MySQL = mo.New() + o.MySQL.UpstreamURL = u.String() + u.Scheme, u.Host = "postgres", pgOrigin + o.Postgres = pgo.New() + o.Postgres.UpstreamURL = u.String() + c.Backends = bo.Lookup{"greptime": o} + }) + originalRelease := harness.releasePorts + harness.releasePorts = func() { release(); originalRelease() } + harness.start(t) + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cancel() + my, err := vtmysql.Connect(ctx, &vtmysql.ConnParams{Host: "127.0.0.1", Port: ports[0], Uname: "grafana_ro", Pass: "trickster-dev-grafana", DbName: "public"}) + require.NoError(t, err) + defer my.Close() + pg, err := pgwireConnect(t, fmt.Sprintf("127.0.0.1:%d", ports[1]), pgwireTarget{ClientUser: "grafana_ro", Database: "public"}, "trickster-dev-grafana") + require.NoError(t, err) + defer pg.Close(context.Background()) + for n := 0; n < 3; n++ { + require.NoError(t, my.GetRawConn().SetDeadline(time.Now().Add(10*time.Second))) + r, err := my.ExecuteFetch("SELECT 42 AS answer", 1, true) + require.NoError(t, err) + require.Len(t, r.Rows, 1) + require.Equal(t, "42", r.Rows[0][0].ToString()) + rows, err := pgwireQuery(t, pg, "SELECT 42 AS answer") + require.NoError(t, err) + require.Equal(t, [][]string{{"42"}}, rows[0].Rows) + req, err := http.NewRequestWithContext(ctx, http.MethodGet, "http://"+harness.BaseAddr+"/greptime/v1/sql?sql="+url.QueryEscape("SELECT 42 AS answer"), nil) + require.NoError(t, err) + req.SetBasicAuth("grafana_ro", "trickster-dev-grafana") + resp, err := http.DefaultClient.Do(req) + require.NoError(t, err) + var body struct { + Code int `json:"code"` + Output []struct { + Records struct { + Rows [][]int64 `json:"rows"` + } `json:"records"` + } `json:"output"` + } + err = json.NewDecoder(resp.Body).Decode(&body) + _ = resp.Body.Close() + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode) + require.Zero(t, body.Code) + require.Len(t, body.Output, 1) + require.Equal(t, [][]int64{{42}}, body.Output[0].Records.Rows) + } +} + +func TestGreptimeMySQLRealServer(t *testing.T) { + if os.Getenv("TRICKSTER_GREPTIMEDB_MYSQL_ACCEPTANCE") != "1" { + t.Skip("set TRICKSTER_GREPTIMEDB_MYSQL_ACCEPTANCE=1 with isolated developer GreptimeDB") + } + address := os.Getenv("GREPTIMEDB_MYSQL_ADDR") + if address == "" { + address = "127.0.0.1:4002" + } + probe, err := net.DialTimeout("tcp", address, time.Second) + require.NoError(t, err, "requested GreptimeDB MySQL acceptance requires a running developer origin") + _ = probe.Close() + host, portText, err := net.SplitHostPort(address) + require.NoError(t, err) + port, err := strconv.Atoi(portText) + require.NoError(t, err) + params := vtmysql.ConnParams{Host: host, Port: port, Uname: "seeder", Pass: "trickster-dev-seed", DbName: "public"} + connect := func(t *testing.T, params vtmysql.ConnParams) *vtmysql.Conn { + t.Helper() + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + c, err := vtmysql.Connect(ctx, ¶ms) + require.NoError(t, err) + t.Cleanup(c.Close) + return c + } + query := func(t *testing.T, c *vtmysql.Conn, sql string) *sqltypes.Result { + t.Helper() + require.NoError(t, c.GetRawConn().SetDeadline(time.Now().Add(15*time.Second))) + r, err := c.ExecuteFetch(sql, 1000, true) + require.NoError(t, err, sql) + return r + } + writer := connect(t, params) + table := fmt.Sprintf("trickster_mysql_%d", time.Now().UnixNano()) + query(t, writer, "CREATE TABLE "+table+" (ts TIMESTAMP(9) TIME INDEX, label STRING, reading BIGINT)") + t.Cleanup(func() { query(t, writer, "DROP TABLE IF EXISTS "+table) }) + query(t, writer, "INSERT INTO "+table+" VALUES "+ + "('2026-01-01 00:00:00.123456789','A',9007199254740993),"+ + "('2026-01-01 00:00:01','a',2),('2026-01-01 00:00:02',NULL,3),"+ + "('2026-01-01 00:01:00','A',4),('2026-01-01 00:01:01','a',5),"+ + "('2026-01-01 00:02:00','A',6),('2026-01-01 00:02:01','a',7),"+ + "('2026-01-01 00:02:02',NULL,8)") + params.Uname, params.Pass = "grafana_ro", "trickster-dev-grafana" + cfg := cacheoptions.New() + cfg.Name = table + cfg.Provider = "memory" + cache := cachemanager.NewCache(cachememory.New(cfg.Name, cfg), cachemanager.CacheOptions{}, cfg) + require.NoError(t, cache.Connect()) + t.Cleanup(func() { require.NoError(t, cache.Close()) }) + server, err := backend.NewProtocolServer(backend.ProtocolConfig{ + BackendName: table, Upstream: params, Engine: greptimedb.MySQLEngine(), Cache: cache, CacheTTL: time.Minute, + DownstreamUsers: map[string]string{"probe": "trickster-dev-probe"}, ConnectTimeout: 5 * time.Second, QueryTimeout: 10 * time.Second, + }) + require.NoError(t, err) + l, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + done := make(chan error, 1) + go func() { done <- server.Serve(l) }() + t.Cleanup(func() { + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + require.NoError(t, server.Shutdown(ctx)) + select { + case err := <-done: + require.NoError(t, err) + case <-ctx.Done(): + t.Fatal(ctx.Err()) + } + }) + newSession := func(t *testing.T) (*vtmysql.Conn, *vtmysql.Conn, func(string)) { + t.Helper() + direct := connect(t, params) + proxy := connect(t, vtmysql.ConnParams{Host: "127.0.0.1", Port: l.Addr().(*net.TCPAddr).Port, Uname: "probe", Pass: "trickster-dev-probe", DbName: "public"}) + return direct, proxy, func(sql string) { + t.Helper() + want, got := query(t, direct, sql), query(t, proxy, sql) + require.Equal(t, want.Fields, got.Fields) + require.Equal(t, want.Rows, got.Rows) + require.Equal(t, want.StatusFlags, got.StatusFlags) + } + } + count := func(mode, status string) float64 { + metric := &dto.Metric{} + require.NoError(t, metrics.SQLQueryCache.WithLabelValues(table, "greptimedb", mode, status).Write(metric)) + return metric.GetCounter().GetValue() + } + t.Run("OPC retains exact timestamp integer and null", func(t *testing.T) { + _, _, agree := newSession(t) + sql := "SELECT ts,label,reading FROM " + table + " ORDER BY ts" + agree(sql) + before := count("object", "hit") + agree(sql) + require.Equal(t, before+1, count("object", "hit")) + }) + t.Run("DPC miss hit and extended range", func(t *testing.T) { + _, _, agree := newSession(t) + for _, bucket := range []string{"DATE_BIN('1m',ts,FROM_UNIXTIME(0))", "DATE_TRUNC('minute',ts)"} { + format := "SELECT " + bucket + " AS time,label,SUM(reading) AS total FROM " + table + " WHERE ts >= FROM_UNIXTIME(1767225600) AND ts < FROM_UNIXTIME(%d) GROUP BY time,label ORDER BY time,label" + short := fmt.Sprintf(format, 1767225720) + long := fmt.Sprintf(format, 1767225780) + agree(short) + before := count("delta", "hit") + agree(short) + require.Equal(t, before+1, count("delta", "hit")) + before = count("delta", "phit") + agree(long) + require.Equal(t, before+1, count("delta", "phit")) + agree(long) + } + }) + t.Run("error leaves connection usable", func(t *testing.T) { + direct, proxy, agree := newSession(t) + _, want := direct.ExecuteFetch("SHOW COUNT(*) WARNINGS", 100, true) + _, got := proxy.ExecuteFetch("SHOW COUNT(*) WARNINGS", 100, true) + require.Error(t, want) + require.Equal(t, want.Error(), got.Error()) + agree("SELECT 17 AS answer") + }) + t.Run("failed session setting conservatively bypasses cache", func(t *testing.T) { + direct, proxy, agree := newSession(t) + sql := "SELECT 23 AS answer" + agree(sql) + _, want := direct.ExecuteFetch("SET time_zone = '+08:00'", 100, true) + _, got := proxy.ExecuteFetch("SET time_zone = '+08:00'", 100, true) + require.Error(t, want) + require.Equal(t, want.Error(), got.Error()) + before := count("object", "hit") + agree("SHOW TIMEZONE") + agree(sql) + require.Equal(t, before, count("object", "hit")) + }) + t.Run("partial buckets align to complete buckets", func(t *testing.T) { + direct, proxy, _ := newSession(t) + statement := "SELECT DATE_BIN('1m',ts,FROM_UNIXTIME(0)) AS time,label,COUNT(*) AS aligned_total FROM " + table + " WHERE ts >= FROM_UNIXTIME(1767225601) AND ts < FROM_UNIXTIME(1767225779) GROUP BY time,label ORDER BY time,label" + aligned := strings.NewReplacer("1767225601", "1767225660", "1767225779", "1767225720").Replace(statement) + want := query(t, direct, aligned) + for _, status := range []string{"kmiss", "hit"} { + before := count("delta", status) + got := query(t, proxy, statement) + require.Equal(t, want.Fields, got.Fields) + require.Equal(t, want.Rows, got.Rows) + require.Equal(t, want.StatusFlags, got.StatusFlags) + require.Equal(t, before+1, count("delta", status)) + } + }) + t.Run("unsupported precision retains original query semantics", func(t *testing.T) { + _, _, agree := newSession(t) + for _, bounds := range []string{ + "ts >= FROM_UNIXTIME(1767225601) AND ts < FROM_UNIXTIME(1767225610)", + "ts >= FROM_UNIXTIME(1767225600) AND ts <= FROM_UNIXTIME(1767225720)", + } { + sql := "SELECT DATE_BIN('1m',ts,FROM_UNIXTIME(0)) AS time,label,COUNT(*) AS total FROM " + table + " WHERE " + bounds + " GROUP BY time,label ORDER BY time,label" + agree(sql) + before := count("object", "hit") + agree(sql) + require.Equal(t, before+1, count("object", "hit")) + } + }) + t.Run("known transaction stubs still bypass caching", func(t *testing.T) { + _, _, agree := newSession(t) + sql := "SELECT 29 AS answer" + agree(sql) + before := count("object", "hit") + agree("BEGIN") + agree(sql) + agree("ROLLBACK") + require.Equal(t, before, count("object", "hit")) + agree(sql) + require.Equal(t, before+1, count("object", "hit")) + }) + t.Run("unknown state disables caching", func(t *testing.T) { + _, _, agree := newSession(t) + agree("SET sql_mode = 'ANSI_QUOTES'") + before := count("object", "hit") + agree("SELECT 19 AS answer") + agree("SELECT 19 AS answer") + require.Equal(t, before, count("object", "hit")) + }) +} diff --git a/integration/greptimedb_prometheus_test.go b/integration/greptimedb_prometheus_test.go new file mode 100644 index 000000000..a72bebd71 --- /dev/null +++ b/integration/greptimedb_prometheus_test.go @@ -0,0 +1,419 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 integration + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "maps" + "math" + "math/big" + "net/http" + "net/url" + "os" + "path/filepath" + "slices" + "strconv" + "strings" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/integration/internal/portutil" + "github.com/trickstercache/trickster/v2/pkg/proxy/headers" + + "github.com/stretchr/testify/require" + "go.yaml.in/yaml/v3" +) + +const greptimePromPath = "/v1/prometheus/api/v1/" + +type greptimePromResponse struct { + Code int `json:"http_status"` + Result string `json:"trickster_result,omitempty"` + Body json.RawMessage `json:"body"` + Text string `json:"text,omitempty"` +} + +type greptimePromFixture struct { + origin, proxy, dir string + client *http.Client + start time.Time + databases []string + sequence int +} + +// The reference uses explicit step-aligned endpoints; the proxy still receives +// the caller's original range. Instant queries and passthrough tests do not use it. +func (f *greptimePromFixture) alignedOrigin(t *testing.T, method, endpoint string, query, form url.Values, hdr http.Header) greptimePromResponse { + t.Helper() + query, form = maps.Clone(query), maps.Clone(form) + if endpoint == greptimePromPath+"query_range" { + values := query + if form != nil { + values = form + } + step, err := time.ParseDuration(values.Get("step") + "s") + require.NoError(t, err) + for _, name := range []string{"start", "end"} { + at, err := time.Parse(time.RFC3339Nano, values.Get(name)) + require.NoError(t, err) + values.Set(name, at.Truncate(step).Format(time.RFC3339Nano)) + } + } + return f.fetch(t, f.origin, method, endpoint, query, form, hdr) +} + +func (f *greptimePromFixture) fetch(t *testing.T, base, method, endpoint string, query, form url.Values, hdr http.Header) greptimePromResponse { + t.Helper() + var body io.Reader + if form != nil { + body = strings.NewReader(form.Encode()) + } + r, err := http.NewRequest(method, base+endpoint+"?"+query.Encode(), body) + require.NoError(t, err) + if hdr != nil { + r.Header = hdr.Clone() + } + if form != nil { + r.Header.Set("Content-Type", "application/x-www-form-urlencoded") + } + r.SetBasicAuth("grafana_ro", "trickster-dev-grafana") + resp, err := f.client.Do(r) + require.NoError(t, err) + defer resp.Body.Close() + raw, err := io.ReadAll(io.LimitReader(resp.Body, 16<<20)) + require.NoError(t, err) + out := greptimePromResponse{Code: resp.StatusCode, Result: resp.Header.Get(headers.NameTricksterResult)} + if json.Valid(raw) { + out.Body = raw + } else { + out.Text = string(raw) + } + f.sequence++ + evidence, err := json.MarshalIndent(map[string]any{ + "test": t.Name(), "method": method, "url": r.URL.String(), "form": form, "response": out, + }, "", " ") + require.NoError(t, err) + require.NoError(t, os.WriteFile(filepath.Join(f.dir, fmt.Sprintf("%03d.json", f.sequence)), evidence, 0o600)) + return out +} + +func (f *greptimePromFixture) sql(t *testing.T, db, stmt string) { + t.Helper() + r, err := http.NewRequest(http.MethodPost, f.origin+"/v1/sql?db="+url.QueryEscape(db), strings.NewReader(url.Values{"sql": {stmt}}.Encode())) + require.NoError(t, err) + r.SetBasicAuth("seeder", "trickster-dev-seed") + r.Header.Set("Content-Type", "application/x-www-form-urlencoded") + resp, err := f.client.Do(r) + require.NoError(t, err) + defer resp.Body.Close() + raw, err := io.ReadAll(io.LimitReader(resp.Body, 1<<20)) + require.NoError(t, err) + require.Equal(t, 200, resp.StatusCode, "%s: %s", stmt, raw) + var result struct { + Code int + Error string + Output []json.RawMessage + } + require.NoError(t, json.Unmarshal(raw, &result)) + require.Zero(t, result.Code, "%s: %s", stmt, raw) + require.Empty(t, result.Error, "%s: %s", stmt, raw) + require.NotEmpty(t, result.Output, "%s: %s", stmt, raw) +} + +// This opt-in test writes only its unique temporary databases. It starts its own +// Trickster daemon on reserved ports and removes fixtures after stopping it. +func newGreptimePromFixture(t *testing.T) *greptimePromFixture { + t.Helper() + if os.Getenv("TRICKSTER_GREPTIMEDB_PROMQL_ACCEPTANCE") != "1" { + t.Skip("set TRICKSTER_GREPTIMEDB_PROMQL_ACCEPTANCE=1 with isolated developer GreptimeDB") + } + origin := os.Getenv("GREPTIMEDB_HTTP_URL") + if origin == "" { + origin = "http://127.0.0.1:4000" + } + root := os.Getenv("GREPTIMEDB_REPORT_DIR") + require.NotEmpty(t, root, "GREPTIMEDB_REPORT_DIR is required") + require.NoError(t, os.MkdirAll(root, 0o700)) + dir, err := os.MkdirTemp(root, "promql-") + require.NoError(t, err) + f := &greptimePromFixture{origin: strings.TrimRight(origin, "/"), dir: dir, client: &http.Client{Timeout: 30 * time.Second}, start: time.Now().UTC().Add(-5 * time.Minute).Truncate(15 * time.Second).Add(125 * time.Millisecond)} + t.Logf("PromQL evidence: %s", dir) + prefix := "trickster_promql_" + strconv.FormatInt(time.Now().UnixNano(), 10) + for i, suffix := range []string{"a", "b", "all"} { + db := prefix + "_" + suffix + f.sql(t, "public", "CREATE DATABASE "+db) + f.databases = append(f.databases, db) + t.Cleanup(func() { f.sql(t, "public", "DROP DATABASE "+db) }) + f.sql(t, db, "CREATE TABLE fixture (greptime_timestamp TIMESTAMP(3) TIME INDEX, greptime_value DOUBLE, host STRING, PRIMARY KEY(host))") + var values []string + for point := -1; point < 20; point++ { + for host, value := range []int{1, 3, 10} { + if (i == 0 && host == 2) || (i == 1 && host != 2) { + continue + } + values = append(values, fmt.Sprintf("(%d, %d, 'h%d')", f.start.Add(time.Duration(point)*15*time.Second).UnixMilli(), value, host)) + } + } + f.sql(t, db, "INSERT INTO fixture VALUES "+strings.Join(values, ",")) + } + ports, release := portutil.Reserve(t, 3) + rewriters := map[string]any{} + backends := map[string]any{} + for _, name := range []string{"direct", "a", "b"} { + backend := map[string]any{ + "provider": "greptimedb", "origin_url": f.origin, "cache_name": "memory", "fast_forward_disable": true, + "healthcheck": map[string]any{"path": "/health", "query": "", "interval": "100ms", "timeout": "2s", "failure_threshold": 1, "recovery_threshold": 1}, + } + if name != "direct" { + index := 0 + if name == "b" { + index = 1 + } + rewriters[name] = map[string]any{"instructions": [][]string{{"param", "set", "db", f.databases[index]}}} + backend["req_rewriter_name"] = name + backend["prometheus"] = map[string]any{"labels": map[string]string{"replica": name}} + } + backends[name] = backend + } + backends["merged"] = map[string]any{"provider": "alb", "alb": map[string]any{"mechanism": "tsm", "output_format": "greptimedb", "pool": []string{"a", "b"}}} + cfg := map[string]any{ + "listeners": map[string]any{"default": map[string]any{"port": ports[0]}, "metrics": map[string]any{"port": ports[1]}, "mgmt": map[string]any{"port": ports[2]}}, + "logging": map[string]any{"log_level": "error"}, "caches": map[string]any{"memory": map[string]string{"provider": "memory"}}, + "request_rewriters": rewriters, "backends": backends, + } + raw, err := yaml.Marshal(cfg) + require.NoError(t, err) + path := filepath.Join(dir, "trickster.yaml") + require.NoError(t, os.WriteFile(path, raw, 0o600)) + release() + runTrickster(t, context.Background(), "-config", path) + waitForTrickster(t, fmt.Sprintf("127.0.0.1:%d", ports[1])) + f.proxy = fmt.Sprintf("http://127.0.0.1:%d", ports[0]) + return f +} + +// Compare numbers rather than their JSON spelling (2 vs 2.0). Rational +// timestamps retain exact precision, including sub-millisecond differences. +// Series order is insignificant; point order, labels and timestamps are exact. +func greptimePromData(t *testing.T, response greptimePromResponse) map[string]any { + t.Helper() + require.Equal(t, 200, response.Code, "%s", response.Body) + var doc map[string]any + d := json.NewDecoder(bytes.NewReader(response.Body)) + d.UseNumber() + require.NoError(t, d.Decode(&doc)) + require.Equal(t, "success", doc["status"], "%s", response.Body) + data, ok := doc["data"].(map[string]any) + require.True(t, ok, "%s", response.Body) + results, ok := data["result"].([]any) + require.True(t, ok, "%s", response.Body) + for _, result := range results { + series := result.(map[string]any) + points, _ := series["values"].([]any) + if point, ok := series["value"].([]any); ok { + points = []any{point} + } + for _, raw := range points { + point := raw.([]any) + require.Len(t, point, 2) + stamp, ok := point[0].(json.Number) + require.True(t, ok) + exact, ok := new(big.Rat).SetString(stamp.String()) + require.True(t, ok) + point[0] = exact.RatString() + value, err := strconv.ParseFloat(point[1].(string), 64) + require.NoError(t, err) + if !math.IsNaN(value) { + point[1] = value + } + } + } + slices.SortFunc(results, func(a, b any) int { + la, _ := json.Marshal(a.(map[string]any)["metric"]) + lb, _ := json.Marshal(b.(map[string]any)["metric"]) + return bytes.Compare(la, lb) + }) + return doc +} + +func TestGreptimePromDataTimestampComparison(t *testing.T) { + data := func(stamp string) map[string]any { + return greptimePromData(t, greptimePromResponse{ + Code: http.StatusOK, + Body: json.RawMessage(`{"status":"success","data":{"resultType":"matrix","result":[{"metric":{},"values":[[` + stamp + `,"14"]]}]}}`), + }) + } + require.Equal(t, data("1790567805"), data("1790567805.0")) + require.Equal(t, data("1790567805.125"), data("1.790567805125e9")) + require.NotEqual(t, data("1790567805.125"), data("1790567805.125000001")) + require.NotEqual(t, data("1790567805"), data("1790567806")) +} + +func TestGreptimeDBPrometheus(t *testing.T) { + f := newGreptimePromFixture(t) + for _, method := range []string{http.MethodGet, http.MethodPost} { + t.Run(method, func(t *testing.T) { + for _, step := range []time.Duration{15 * time.Second, 500 * time.Millisecond} { + t.Run("cache_"+step.String(), func(t *testing.T) { + v := url.Values{"query": {"sum(fixture)"}, "db": {f.databases[2]}, "start": {f.start.Format(time.RFC3339Nano)}, "step": {strconv.FormatFloat(step.Seconds(), 'f', -1, 64)}} + for i, status := range []string{"kmiss", "hit", "phit", "hit"} { + points := 3 + if i >= 2 { + points = 5 + } + v.Set("end", f.start.Add(time.Duration(points-1)*step).Format(time.RFC3339Nano)) + query, form := v, url.Values(nil) + if method == http.MethodPost { + query = url.Values{"db": v["db"]} + form = v + } + origin := f.alignedOrigin(t, method, greptimePromPath+"query_range", query, form, nil) + proxy := f.fetch(t, f.proxy+"/direct", method, greptimePromPath+"query_range", query, form, nil) + require.Equal(t, greptimePromData(t, origin), greptimePromData(t, proxy)) + engine, got := headers.ParseResultEngineStatus(proxy.Result) + require.Equal(t, "DeltaProxyCache", engine) + require.Equal(t, status, got) + var matrix struct { + Data struct{ Result []struct{ Values [][]any } } + } + require.NoError(t, json.Unmarshal(proxy.Body, &matrix)) + require.Len(t, matrix.Data.Result, 1) + require.Len(t, matrix.Data.Result[0].Values, points) + for n, point := range matrix.Data.Result[0].Values { + require.Equal(t, float64(f.start.Truncate(step).Add(time.Duration(n)*step).UnixMilli())/1000, point[0]) + value, err := strconv.ParseFloat(point[1].(string), 64) + require.NoError(t, err) + require.Equal(t, float64(14), value) + } + } + }) + } + for _, endpoint := range []string{"query", "query_range"} { + for _, expression := range []string{"sum(fixture)", "count(fixture)", "avg(fixture)", "min(fixture)", "max(fixture)", "group(fixture)", "quantile(0.5, fixture)", "quantile by (host) (0.5, fixture)", "topk(2, fixture)", "bottomk(2, fixture)", "sort(count(fixture) or vector(0))", "sum(fixture) / 2", "avg by (host) (fixture)"} { + t.Run("merge_"+endpoint+"_"+expression, func(t *testing.T) { + v := url.Values{"query": {expression}, "db": {f.databases[2]}} + if endpoint == "query" { + v.Set("time", f.start.Add(30*time.Second).Format(time.RFC3339Nano)) + } else { + v.Set("start", f.start.Format(time.RFC3339Nano)) + v.Set("end", f.start.Add(30*time.Second).Format(time.RFC3339Nano)) + v.Set("step", "15") + } + query, form := v, url.Values(nil) + if method == http.MethodPost { + query = url.Values{"db": v["db"]} + form = v + } + origin := f.alignedOrigin(t, method, greptimePromPath+endpoint, query, form, nil) + proxy := f.fetch(t, f.proxy+"/merged", method, greptimePromPath+endpoint, query, form, nil) + want := greptimePromData(t, origin) + if endpoint == "query_range" && strings.HasPrefix(expression, "sort(") { + // Shared TSM finalization adds this exact advisory. Other + // warnings remain part of the strict comparison. + want["warnings"] = []any{"PromQL warning: sort is ineffective for range queries since results are always ordered by labels"} + } + require.Equal(t, want, greptimePromData(t, proxy), "origin=%s proxy=%s", origin.Body, proxy.Body) + }) + } + } + }) + } + for _, useHeader := range []bool{false, true} { + for i, db := range f.databases[:2] { + t.Run(fmt.Sprintf("database_%t_%d", useHeader, i), func(t *testing.T) { + v := url.Values{"query": {"sum(fixture)"}, "start": {f.start.Format(time.RFC3339Nano)}, "end": {f.start.Add(30 * time.Second).Format(time.RFC3339Nano)}, "step": {"15"}} + hdr := make(http.Header) + if useHeader { + hdr.Set("X-Greptime-Db-Name", db) + } else { + v.Set("db", db) + } + origin := f.alignedOrigin(t, http.MethodGet, greptimePromPath+"query_range", v, nil, hdr) + for range 2 { + proxy := f.fetch(t, f.proxy+"/direct", http.MethodGet, greptimePromPath+"query_range", v, nil, hdr) + require.Equal(t, greptimePromData(t, origin), greptimePromData(t, proxy)) + } + }) + } + } + t.Run("url_form_precedence", func(t *testing.T) { + v := url.Values{"query": {"avg(fixture)"}, "db": {f.databases[2]}, "time": {f.start.Add(30 * time.Second).Format(time.RFC3339Nano)}} + form := maps.Clone(v) + form.Set("query", "min(fixture)") + for _, backend := range []string{"direct", "merged"} { + origin := f.fetch(t, f.origin, http.MethodPost, greptimePromPath+"query", v, form, nil) + proxy := f.fetch(t, f.proxy+"/"+backend, http.MethodPost, greptimePromPath+"query", v, form, nil) + require.Equal(t, greptimePromData(t, origin), greptimePromData(t, proxy)) + } + }) + for _, lookback := range []string{"1s", "1m"} { + t.Run("lookback_"+lookback, func(t *testing.T) { + v := url.Values{"db": {f.databases[2]}, "query": {"fixture"}, "start": {f.start.Add(5 * time.Second).Format(time.RFC3339Nano)}, "end": {f.start.Add(35 * time.Second).Format(time.RFC3339Nano)}, "step": {"15"}, "lookback": {lookback}} + origin := f.alignedOrigin(t, "GET", greptimePromPath+"query_range", v, nil, nil) + for range 2 { + proxy := f.fetch(t, f.proxy+"/direct", "GET", greptimePromPath+"query_range", v, nil, nil) + require.Equal(t, greptimePromData(t, origin), greptimePromData(t, proxy)) + } + }) + } + for _, endpoint := range []string{"series", "labels", "label/host/values"} { + t.Run(endpoint, func(t *testing.T) { + v := url.Values{"db": {f.databases[2]}, "match[]": {"fixture"}, "start": {f.start.Format(time.RFC3339Nano)}, "end": {f.start.Add(30 * time.Second).Format(time.RFC3339Nano)}} + origin := f.fetch(t, f.origin, "GET", greptimePromPath+endpoint, v, nil, nil) + proxy := f.fetch(t, f.proxy+"/direct", "GET", greptimePromPath+endpoint, v, nil, nil) + require.Equal(t, 200, origin.Code) + require.Equal(t, origin.Code, proxy.Code) + var a, b struct { + Status string + Data []any + } + require.NoError(t, json.Unmarshal(origin.Body, &a)) + require.NoError(t, json.Unmarshal(proxy.Body, &b)) + require.Equal(t, "success", a.Status) + require.Equal(t, a.Status, b.Status) + require.ElementsMatch(t, a.Data, b.Data) + }) + } + for _, query := range []string{`fixture{host="missing"}`, "sum("} { + t.Run("empty_or_error_"+query, func(t *testing.T) { + v := url.Values{"db": {f.databases[2]}, "query": {query}, "start": {f.start.Format(time.RFC3339Nano)}, "end": {f.start.Add(30 * time.Second).Format(time.RFC3339Nano)}, "step": {"15"}} + origin := f.fetch(t, f.origin, "GET", greptimePromPath+"query_range", v, nil, nil) + proxy := f.fetch(t, f.proxy+"/direct", "GET", greptimePromPath+"query_range", v, nil, nil) + require.Equal(t, origin.Code, proxy.Code) + require.JSONEq(t, string(origin.Body), string(proxy.Body)) + }) + } + t.Run("count_values_origin_limitation", func(t *testing.T) { + for _, endpoint := range []string{"query", "query_range"} { + v := url.Values{"db": {f.databases[2]}, "query": {`count_values("value", fixture)`}, "start": {f.start.Format(time.RFC3339Nano)}, "end": {f.start.Add(30 * time.Second).Format(time.RFC3339Nano)}, "step": {"15"}, "time": {f.start.Format(time.RFC3339Nano)}} + origin := f.fetch(t, f.origin, "GET", greptimePromPath+endpoint, v, nil, nil) + for range 2 { + proxy := f.fetch(t, f.proxy+"/direct", "GET", greptimePromPath+endpoint, v, nil, nil) + require.Equal(t, origin.Code, proxy.Code) + require.JSONEq(t, string(origin.Body), string(proxy.Body)) + engine, _ := headers.ParseResultEngineStatus(proxy.Result) + require.Equal(t, "HTTPProxy", engine) + } + merged := f.fetch(t, f.proxy+"/merged", "GET", greptimePromPath+endpoint, v, nil, nil) + require.Equal(t, http.StatusBadGateway, merged.Code, "must not merge malformed upstream series") + } + }) +} diff --git a/integration/greptimedb_session_test.go b/integration/greptimedb_session_test.go new file mode 100644 index 000000000..70dc89d8d --- /dev/null +++ b/integration/greptimedb_session_test.go @@ -0,0 +1,113 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 integration + +import ( + "context" + "fmt" + "net" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/backends/providers" + + "github.com/stretchr/testify/require" +) + +// TestGreptimePGSessionSettings uses only SELECT/SHOW and session-local settings. +// It never changes fixture tables or account permissions. +func TestGreptimePGSessionSettings(t *testing.T) { + var target pgwireTarget + for _, candidate := range pgwireTargets() { + if candidate.Provider == providers.GreptimeDB { + target = candidate + break + } + } + require.NotEmpty(t, target.OriginAddr) + probe, err := net.DialTimeout("tcp", target.OriginAddr, time.Second) + if err != nil { + t.Skipf("developer GreptimeDB is unavailable: %v", err) + } + _ = probe.Close() + harness, address := pgwireHarness(t, target) + harness.start(t) + direct, err := pgwireConnect(t, target.OriginAddr, target, target.ClientPassword) + require.NoError(t, err) + defer direct.Close(context.Background()) + proxy, err := pgwireConnect(t, address, target, target.ClientPassword) + require.NoError(t, err) + defer proxy.Close(context.Background()) + start := time.Now().UTC().Add(-48 * time.Hour).Truncate(time.Hour) + query := func(template string) string { + return fmt.Sprintf(template, start.Format(time.RFC3339), start.Add(2*time.Hour).Format(time.RFC3339)) + } + agree := func(sql string) { + t.Helper() + want, err := pgwireQuery(t, direct, sql) + require.NoError(t, err, sql) + got, err := pgwireQuery(t, proxy, sql) + require.NoError(t, err, sql) + require.Equal(t, want, got, sql) + } + stampSQL, epochSQL := query(target.DeltaSQLs[0]), query(target.DeltaSQLs[1]) + agree(stampSQL) + agree(stampSQL) + t.Run("failed SET preserves cache identity", func(t *testing.T) { + _, wantErr := pgwireQuery(t, direct, "SET time_zone = 'not/a/timezone'") + _, gotErr := pgwireQuery(t, proxy, "SET time_zone = 'not/a/timezone'") + require.Error(t, wantErr) + require.Equal(t, pgwireSQLState(wantErr), pgwireSQLState(gotErr)) + before := pgwireCacheCount(t, harness.MetricsAddr, target.Dialect, "delta", "hit") + agree(stampSQL) + require.Equal(t, before+1, pgwireCacheCount(t, harness.MetricsAddr, target.Dialect, "delta", "hit")) + }) + t.Run("no-op PostgreSQL settings keep lossless epochs cacheable", func(t *testing.T) { + agree(epochSQL) + before := pgwireCacheCount(t, harness.MetricsAddr, target.Dialect, "delta", "hit") + agree("SET extra_float_digits = -14") + agree(epochSQL) + agree("SET standard_conforming_strings = off") + agree(epochSQL) + require.Equal(t, before+2, pgwireCacheCount(t, harness.MetricsAddr, target.Dialect, "delta", "hit")) + }) + t.Run("date style changes do not reuse timestamp bytes", func(t *testing.T) { + agree("SET DateStyle = 'SQL, DMY'") + agree(stampSQL) + agree(stampSQL) + agree("SET DateStyle = 'ISO, MDY'") + agree(stampSQL) + agree(stampSQL) + }) + t.Run("LOCAL alias persists but partitions the cache", func(t *testing.T) { + before := pgwireCacheCount(t, harness.MetricsAddr, target.Dialect, "delta", "kmiss") + agree("SET LOCAL time_zone = 'Asia/Kolkata'") + agree("SHOW TIMEZONE") + agree(epochSQL) + require.Equal(t, before+1, pgwireCacheCount(t, harness.MetricsAddr, target.Dialect, "delta", "kmiss")) + before = pgwireCacheCount(t, harness.MetricsAddr, target.Dialect, "delta", "hit") + agree(epochSQL) + require.Equal(t, before+1, pgwireCacheCount(t, harness.MetricsAddr, target.Dialect, "delta", "hit")) + }) + t.Run("an unknown setting disables caching", func(t *testing.T) { + before := pgwireCacheCount(t, harness.MetricsAddr, target.Dialect, "delta", "hit") + agree("SET trickster_phase5_unknown = 1") + agree(epochSQL) + agree(epochSQL) + require.Equal(t, before, pgwireCacheCount(t, harness.MetricsAddr, target.Dialect, "delta", "hit")) + }) +} diff --git a/integration/harness_test.go b/integration/harness_test.go index 9c1be6df0..9ee99f058 100644 --- a/integration/harness_test.go +++ b/integration/harness_test.go @@ -295,6 +295,8 @@ func writeTestConfig(t *testing.T, configPath string, } } } + // GreptimeDB's second native endpoint is likewise opt-in in this harness. + delete(c.Listeners, "greptimedb-mysql") // The dev config binds its PostgreSQL wire-protocol listeners to fixed // ports; drop them and the backends they serve, which need such a listener. // Tests that want one add it back on a reserved port through mods. diff --git a/integration/pgwire_conformance_test.go b/integration/pgwire_conformance_test.go index 7650e9a97..282d5c22c 100644 --- a/integration/pgwire_conformance_test.go +++ b/integration/pgwire_conformance_test.go @@ -21,12 +21,14 @@ import ( "errors" "fmt" "net" + "net/http" + "net/http/httptest" "net/url" - "strconv" "strings" "testing" "time" + "github.com/trickstercache/trickster/v2/integration/internal/metricsutil" "github.com/trickstercache/trickster/v2/integration/internal/portutil" bo "github.com/trickstercache/trickster/v2/pkg/backends/options" "github.com/trickstercache/trickster/v2/pkg/backends/providers" @@ -53,6 +55,7 @@ const ( type pgwireTarget struct { Name string Provider string + Dialect string OriginAddr string Database string OriginUser string @@ -71,6 +74,7 @@ type pgwireTarget struct { DeltaSQLs []string WeekSQL string ZoneSQL string + ZoneChangesResults bool SupportsCancel bool SupportsTransactions bool } @@ -79,7 +83,7 @@ func pgwireTargets() []pgwireTarget { // The suite is the same for every engine the postgres listener can front; // only these facts differ, so serving another engine means adding an entry. return []pgwireTarget{{ - Name: "timescaledb", Provider: providers.TimescaleDB, OriginAddr: "127.0.0.1:5432", + Name: "timescaledb", Provider: providers.TimescaleDB, Dialect: providers.Postgres, OriginAddr: "127.0.0.1:5432", Database: "trickster", OriginUser: "trickster", OriginPassword: "trickster-dev-upstream", ClientUser: "grafana_ro", ClientPassword: "trickster-dev-grafana", ScalarSQL: "SELECT 42::int4 AS i, 'text'::text AS t, 1.50::numeric AS n, NULL::text AS z, " + @@ -106,11 +110,49 @@ func pgwireTargets() []pgwireTarget { // week buckets sit on the engine's own grid, which is not the Unix epoch's WeekSQL: "SELECT time_bucket('7 days', pickup_datetime) AS time, count(*) FROM trips " + "WHERE pickup_datetime >= '%s' AND pickup_datetime < '%s' GROUP BY 1 ORDER BY 1", - ZoneSQL: "SET TIME ZONE 'Asia/Kolkata'", + ZoneSQL: "SET TIME ZONE 'Asia/Kolkata'", ZoneChangesResults: true, SupportsCancel: true, SupportsTransactions: true, + }, { + Name: "greptimedb", Provider: providers.GreptimeDB, Dialect: providers.GreptimeDB, OriginAddr: "127.0.0.1:4003", + // GreptimeDB's read-only users cannot SET session variables. This fixture + // account is used only for conformance; the Grafana backend remains read-only. + Database: "public", OriginUser: "seeder", OriginPassword: "trickster-dev-seed", + ClientUser: "seeder", ClientPassword: "trickster-dev-seed", + ScalarSQL: "SELECT CAST(42 AS INT) AS i, 'text' AS t, CAST(1.50 AS DOUBLE) AS n, " + + "CAST(NULL AS STRING) AS z, TIMESTAMP '2026-01-02 03:04:05' AS ts, true AS b", + LargeSQL: "SELECT pickup_epoch, cab_type, passenger_count FROM trips " + + "ORDER BY pickup_epoch, cab_type, passenger_count LIMIT 50000", LargeRows: 50000, + MissingSQL: "SELECT * FROM __missing_pgwire_conformance_table", + SetShowName: "TimeZone", SetShowSQL: "SET time_zone = 'Asia/Kolkata'", ShowSQL: "SHOW TIMEZONE", + ObjectSQL: "SELECT cab_type, count(*) FROM trips GROUP BY 1 ORDER BY 1", + DeltaSQLs: []string{ + "SELECT date_bin(INTERVAL '5 minutes', pickup_datetime) AS time, cab_type, count(*) AS trips " + + "FROM trips WHERE pickup_datetime >= '%s' AND pickup_datetime < '%s' GROUP BY 1, 2 ORDER BY 1, 2", + "SELECT floor(extract(epoch FROM pickup_datetime)/900)*900 AS time, count(*) AS trips " + + "FROM trips WHERE pickup_datetime >= '%s' AND pickup_datetime < '%s' GROUP BY 1 ORDER BY 1", + "SELECT date_bin('5m', pickup_datetime) AS time, count(*) AS trips " + + "FROM trips WHERE pickup_datetime >= '%s' AND pickup_datetime < '%s' GROUP BY 1 ORDER BY 1 DESC", + "SELECT date_bin(INTERVAL '15 minutes', \"bucket\") AS time, sum(trips) AS trips " + + "FROM trips_15m WHERE \"bucket\" >= '%s' AND \"bucket\" < '%s' GROUP BY 1 ORDER BY 1", + }, + WeekSQL: "SELECT date_trunc('week', pickup_datetime) AS time, count(*) FROM trips " + + "WHERE pickup_datetime >= '%s' AND pickup_datetime < '%s' GROUP BY 1 ORDER BY 1", + ZoneSQL: "SET LOCAL time_zone = 'Asia/Kolkata'", + // The tested origin stubs transaction status and CancelRequest; neither + // capability is usable. + SupportsCancel: false, SupportsTransactions: false, }} } +func pgwireRequireSQL(t *testing.T, scenario string, statements ...string) { + t.Helper() + for _, sql := range statements { + if sql == "" { + t.Skipf("the engine has no %s SQL configured", scenario) + } + } +} + func TestPGWireConformance(t *testing.T) { for _, target := range pgwireTargets() { t.Run(target.Name, func(t *testing.T) { @@ -275,6 +317,7 @@ func runPGWireConformance(t *testing.T, target pgwireTarget, proxyAddr, metricsA }) t.Run("an error keeps its SQLSTATE and the session survives", func(t *testing.T) { + pgwireRequireSQL(t, "error recovery", target.MissingSQL) _, _, wantErr, gotErr := both(target.MissingSQL) require.Error(t, wantErr) require.Equal(t, pgwireSQLState(wantErr), pgwireSQLState(gotErr)) @@ -288,11 +331,14 @@ func runPGWireConformance(t *testing.T, target pgwireTarget, proxyAddr, metricsA if !target.SupportsTransactions { t.Skip("the engine has no transactions") } + pgwireRequireSQL(t, "transaction error", target.MissingSQL) for _, step := range []struct { sql string status byte }{ - {"BEGIN", pgwireTxOpen}, {"SELECT 1", pgwireTxOpen}, {target.MissingSQL, pgwireTxFailed}, + {"BEGIN", pgwireTxOpen}, + {"SELECT 1", pgwireTxOpen}, + {target.MissingSQL, pgwireTxFailed}, {"ROLLBACK", pgwireTxIdle}, } { _, _ = pgwireQuery(t, direct, step.sql) @@ -303,6 +349,7 @@ func runPGWireConformance(t *testing.T, target pgwireTarget, proxyAddr, metricsA }) t.Run("session settings reach the origin and are reported back", func(t *testing.T) { + pgwireRequireSQL(t, "session settings", target.SetShowSQL, target.ShowSQL) _, err := pgwireQuery(t, proxied, target.SetShowSQL) require.NoError(t, err) _, err = pgwireQuery(t, direct, target.SetShowSQL) @@ -315,6 +362,7 @@ func runPGWireConformance(t *testing.T, target pgwireTarget, proxyAddr, metricsA }) t.Run("a large result arrives whole", func(t *testing.T) { + pgwireRequireSQL(t, "large result", target.LargeSQL) want, got, wantErr, gotErr := both(target.LargeSQL) require.NoError(t, wantErr) require.NoError(t, gotErr) @@ -326,6 +374,7 @@ func runPGWireConformance(t *testing.T, target pgwireTarget, proxyAddr, metricsA if !target.SupportsCancel { t.Skip("the engine does not implement cancel requests") } + pgwireRequireSQL(t, "cancellation", target.SlowSQL) failed := make(chan error, 1) go func() { _, err := pgwireQuery(t, proxied, target.SlowSQL) @@ -346,6 +395,10 @@ func runPGWireConformance(t *testing.T, target pgwireTarget, proxyAddr, metricsA }) t.Run("cached results agree with the origin", func(t *testing.T) { + pgwireRequireSQL(t, "object cache", target.ObjectSQL) + if len(target.DeltaSQLs) == 0 { + t.Skip("the engine has no delta-cache SQL configured") + } // a fresh session: the one above changed its time zone, and with it its cache identity cached, err := pgwireConnect(t, proxyAddr, target, target.ClientPassword) require.NoError(t, err) @@ -385,6 +438,10 @@ func runPGWireConformance(t *testing.T, target pgwireTarget, proxyAddr, metricsA }) t.Run("a cached range answers its sub-ranges and later ranges fetch only what is new", func(t *testing.T) { + if len(target.DeltaSQLs) == 0 { + t.Skip("the engine has no delta-cache SQL configured") + } + pgwireRequireSQL(t, "delta cache", target.DeltaSQLs[0]) cached, err := pgwireConnect(t, proxyAddr, target, target.ClientPassword) require.NoError(t, err) defer cached.Close(context.Background()) @@ -401,18 +458,19 @@ func runPGWireConformance(t *testing.T, target pgwireTarget, proxyAddr, metricsA status string // the statement's key is range-independent and already cached elsewhere on the axis }{{0, 6, "rmiss"}, {2, 4, "hit"}, {0, 6, "hit"}, {4, 9, "phit"}, {1, 8, "hit"}} { - before := pgwireCacheCount(t, metricsAddr, "delta", step.status) + before := pgwireCacheCount(t, metricsAddr, target.Dialect, "delta", step.status) want, err := pgwireQuery(t, fresh, statement(step.from, step.to)) require.NoError(t, err) got, err := pgwireQuery(t, cached, statement(step.from, step.to)) require.NoError(t, err) require.Equal(t, want, got, "hours %d to %d", step.from, step.to) - require.Equal(t, before+1, pgwireCacheCount(t, metricsAddr, "delta", step.status), + require.Equal(t, before+1, pgwireCacheCount(t, metricsAddr, target.Dialect, "delta", step.status), "hours %d to %d should be a %s", step.from, step.to, step.status) } }) t.Run("week buckets land on the origin's grid", func(t *testing.T) { + pgwireRequireSQL(t, "week buckets", target.WeekSQL) cached, err := pgwireConnect(t, proxyAddr, target, target.ClientPassword) require.NoError(t, err) defer cached.Close(context.Background()) @@ -425,7 +483,7 @@ func runPGWireConformance(t *testing.T, target pgwireTarget, proxyAddr, metricsA upper = upper.Add(-24 * time.Hour) } sql := fmt.Sprintf(target.WeekSQL, upper.Add(-21*24*time.Hour).Format(time.RFC3339), upper.Format(time.RFC3339)) - before := pgwireCacheCount(t, metricsAddr, "delta", "hit") + before := pgwireCacheCount(t, metricsAddr, target.Dialect, "delta", "hit") want, err := pgwireQuery(t, fresh, sql) require.NoError(t, err) require.Len(t, want[0].Rows, 3) @@ -435,10 +493,15 @@ func runPGWireConformance(t *testing.T, target pgwireTarget, proxyAddr, metricsA require.Equal(t, want, got) } // a bucket off the planned grid would have been refused and the statement relayed uncached - require.Equal(t, before+1, pgwireCacheCount(t, metricsAddr, "delta", "hit")) + require.Equal(t, before+1, pgwireCacheCount(t, metricsAddr, target.Dialect, "delta", "hit")) }) t.Run("sessions with different identities do not share answers", func(t *testing.T) { + pgwireRequireSQL(t, "session time zone", target.ZoneSQL) + if len(target.DeltaSQLs) < 2 { + t.Skip("the engine has no session-aware delta-cache SQL configured") + } + pgwireRequireSQL(t, "session cache identity", target.DeltaSQLs[1]) start := time.Now().UTC().Add(-120 * time.Hour).Truncate(time.Hour) // the second statement stays delta-cacheable under any session zone sql := fmt.Sprintf(target.DeltaSQLs[1], start.Format(time.RFC3339), start.Add(2*time.Hour).Format(time.RFC3339)) @@ -456,25 +519,29 @@ func runPGWireConformance(t *testing.T, target pgwireTarget, proxyAddr, metricsA require.NoError(t, err) return out } - misses := func() float64 { return pgwireCacheCount(t, metricsAddr, "delta", "kmiss") } + misses := func() float64 { return pgwireCacheCount(t, metricsAddr, target.Dialect, "delta", "kmiss") } before := misses() utc := session(proxyAddr, target.ClientUser) require.Equal(t, session(target.OriginAddr, target.ClientUser), utc) // this identity already holds the statement over other ranges, so it is no key miss require.Equal(t, before, misses()) - // a time zone changes how every timestamp is rendered, so it is another cache entry + // The session zone partitions the cache even for engines that render naive UTC timestamps. zoned := session(proxyAddr, target.ClientUser, target.ZoneSQL) require.Equal(t, session(target.OriginAddr, target.ClientUser, target.ZoneSQL), zoned) - require.NotEqual(t, utc, zoned) + if target.ZoneChangesResults { + require.NotEqual(t, utc, zoned) + } else { + require.Equal(t, utc, zoned) + } require.Equal(t, before+1, misses()) // and so does the client's user, which row-level security may depend on require.Equal(t, utc, session(proxyAddr, pgwireSecondUser)) require.Equal(t, before+2, misses()) // each identity is served from its own entry afterwards - hits := pgwireCacheCount(t, metricsAddr, "delta", "hit") + hits := pgwireCacheCount(t, metricsAddr, target.Dialect, "delta", "hit") require.Equal(t, utc, session(proxyAddr, target.ClientUser)) require.Equal(t, zoned, session(proxyAddr, target.ClientUser, target.ZoneSQL)) - require.Equal(t, hits+2, pgwireCacheCount(t, metricsAddr, "delta", "hit")) + require.Equal(t, hits+2, pgwireCacheCount(t, metricsAddr, target.Dialect, "delta", "hit")) _, body := getBody(t, "http://"+metricsAddr+"/metrics") require.NotContains(t, body, `trickster_sql_query_rewrite_failures_total{backend_name="`+pgwireConformanceName+`"`) @@ -482,23 +549,31 @@ func runPGWireConformance(t *testing.T, target pgwireTarget, proxyAddr, metricsA t.Run("relayed statements are classified", func(t *testing.T) { _, body := getBody(t, "http://"+metricsAddr+"/metrics") - require.Contains(t, body, `trickster_sql_query_analysis_total{backend_name="`+pgwireConformanceName+`"`) + if target.ObjectSQL != "" || len(target.DeltaSQLs) != 0 { + require.Contains(t, body, `trickster_sql_query_analysis_total{backend_name="`+pgwireConformanceName+`"`) + } require.Contains(t, body, `trickster_proxy_requests_total{backend_name="`+pgwireConformanceName+`"`) }) } -func pgwireCacheCount(t *testing.T, metricsAddr, mode, status string) float64 { +func pgwireCacheCount(t *testing.T, metricsAddr, dialect, mode, status string) float64 { t.Helper() - _, body := getBody(t, "http://"+metricsAddr+"/metrics") - prefix := `trickster_sql_query_cache_total{backend_name="` + pgwireConformanceName + `",cache_mode="` + mode + - `",cache_status="` + status + `"` - for _, line := range strings.Split(body, "\n") { - if !strings.HasPrefix(line, prefix) { - continue - } - value, err := strconv.ParseFloat(line[strings.LastIndexByte(line, ' ')+1:], 64) - require.NoError(t, err) - return value - } - return 0 + metrics := metricsutil.ScrapeURL(t, "http://"+metricsAddr+"/metrics", nil) + return metrics[metricsutil.Key("trickster_sql_query_cache_total", map[string]string{ + "backend_name": pgwireConformanceName, "dialect": dialect, "cache_mode": mode, "cache_status": status, + })] +} + +func TestPGWireCacheCountSeparatesDialects(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + _, _ = fmt.Fprintln(w, `trickster_sql_query_cache_total{backend_name="pgwire-conformance",cache_mode="delta",cache_status="hit",dialect="greptimedb"} 7`) + _, _ = fmt.Fprintln(w, `trickster_sql_query_cache_total{dialect="postgres",cache_status="hit",cache_mode="delta",backend_name="pgwire-conformance"} 3`) + _, _ = fmt.Fprintln(w, `trickster_sql_query_cache_total{backend_name="other",cache_mode="delta",cache_status="hit",dialect="postgres"} 9`) + _, _ = fmt.Fprintln(w, `trickster_sql_query_cache_total{backend_name="pgwire-conformance",cache_mode="object",cache_status="hit",dialect="postgres"} 11`) + })) + defer server.Close() + address := strings.TrimPrefix(server.URL, "http://") + require.Equal(t, float64(7), pgwireCacheCount(t, address, providers.GreptimeDB, "delta", "hit")) + require.Equal(t, float64(3), pgwireCacheCount(t, address, providers.Postgres, "delta", "hit")) + require.Zero(t, pgwireCacheCount(t, address, providers.Postgres, "delta", "kmiss")) } diff --git a/pkg/backends/clickhouse/clickhouse_test.go b/pkg/backends/clickhouse/clickhouse_test.go index 2a91b18eb..9d36a8543 100644 --- a/pkg/backends/clickhouse/clickhouse_test.go +++ b/pkg/backends/clickhouse/clickhouse_test.go @@ -199,7 +199,7 @@ func TestNativeListenerAdapterValidation(t *testing.T) { if NativeListenerAdapter().Protocol() != listenerconfig.ProtocolClickHouse { t.Fatal("exported adapter has wrong protocol") } - if a.Protocol() != listenerconfig.ProtocolClickHouse || !a.SupportsHTTP() || a.Configured(nil) { + if a.Protocol() != listenerconfig.ProtocolClickHouse || !a.SupportsHTTP(providers.ClickHouse) || a.Configured(nil) { t.Fatal("unexpected ClickHouse adapter capabilities") } if err := a.ValidateListener(nil); err == nil { diff --git a/pkg/backends/clickhouse/native_listener.go b/pkg/backends/clickhouse/native_listener.go index 47f94548b..bc988b58c 100644 --- a/pkg/backends/clickhouse/native_listener.go +++ b/pkg/backends/clickhouse/native_listener.go @@ -41,7 +41,7 @@ type nativeListenerAdapter struct{} // NativeListenerAdapter returns ClickHouse's shared native-listener adapter. func NativeListenerAdapter() native.Adapter { return nativeListenerAdapter{} } -func (nativeListenerAdapter) SupportsHTTP() bool { return true } +func (nativeListenerAdapter) SupportsHTTP(string) bool { return true } func (nativeListenerAdapter) Protocol() string { return listenerconfig.ProtocolClickHouse } diff --git a/pkg/backends/greptimedb/alignment_test.go b/pkg/backends/greptimedb/alignment_test.go new file mode 100644 index 000000000..fcb7a4ebd --- /dev/null +++ b/pkg/backends/greptimedb/alignment_test.go @@ -0,0 +1,113 @@ +/* + * Copyright 2026 The Trickster Authors + * 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 greptimedb + +import ( + "fmt" + "net/http/httptest" + "net/url" + "strings" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/backends/mysql" + "github.com/trickstercache/trickster/v2/pkg/parsing/sqlanalyzer" + "github.com/trickstercache/trickster/v2/pkg/proxy/pgwire" + "github.com/trickstercache/trickster/v2/pkg/timeseries" +) + +func TestSQLCompleteBucketAlignment(t *testing.T) { + pg := analyzer.ForSession(pgwire.SessionView{UTC: true}) + for _, tc := range []struct { + name, lower, upper, operator string + start, end int + }{ + {"lower", "00:00:01", "01:00:00", "<", 15, 45}, + {"upper", "00:00:00", "00:59:59", "<", 0, 30}, + {"both", "00:00:01", "00:59:59", "<", 15, 30}, + {"inclusive", "00:00:00", "01:00:00", "<=", 0, 45}, + } { + statement := strings.NewReplacer("00:00:00Z", tc.lower+"Z", "01:00:00Z", tc.upper+"Z", "ts <", "ts "+tc.operator).Replace(httpSQL) + for _, protocol := range []string{"GET", "POST", "PGWire"} { + t.Run(tc.name+"/"+protocol, func(t *testing.T) { + var trq *timeseries.TimeRangeQuery + if protocol == "PGWire" { + a := pg.Analyze(statement, time.Time{}) + if a.Mode != sqlanalyzer.CacheModeDelta { + t.Fatalf("not delta: %+v", a) + } + trq = sqlanalyzer.NewTimeRangeQuery(statement) + a.Plan.ApplyToQuery(trq) + trq.Extent = a.Plan.RequestExtent(time.Time{}) + } else { + values := url.Values{"sql": {statement}}.Encode() + r := httptest.NewRequest(protocol, "/v1/sql?"+values, nil) + if protocol == "POST" { + r = httptest.NewRequest(protocol, "/v1/sql", strings.NewReader(values)) + r.Header.Set("Content-Type", "application/x-www-form-urlencoded") + } + var err error + trq, _, _, err = (&Client{}).ParseTimeRangeQuery(r) + if err != nil { + t.Fatal(err) + } + } + base := time.Date(2024, 1, 1, 0, 0, 0, 0, time.UTC) + want := timeseries.Extent{Start: base.Add(time.Duration(tc.start) * time.Minute), End: base.Add(time.Duration(tc.end) * time.Minute)} + if !trq.Extent.Start.Equal(want.Start) || !trq.Extent.End.Equal(want.End) { + t.Fatalf("extent=%+v want=%+v", trq.Extent, want) + } + plan := trq.ParsedQuery.(*sqlanalyzer.QueryPlan) + rendered, err := plan.RenderExtent(want) + if err != nil { + t.Fatal(err) + } + a := pg.Analyze(rendered, time.Time{}) + if a.Mode != sqlanalyzer.CacheModeDelta || a.Plan.CanonicalSQL != plan.CanonicalSQL || a.Plan.DropsPartialBuckets { + t.Fatalf("rendered range is not complete buckets: %s %+v", rendered, a) + } + }) + } + } +} + +func TestMySQLCompleteBucketAlignment(t *testing.T) { + a := MySQLEngine().Analyzer(mysql.SessionView{TimeZone: "UTC"}) + for _, tc := range []struct{ lower, upper, start, end int64 }{ + {1767225601, 1767225720, 1767225660, 1767225660}, + {1767225600, 1767225719, 1767225600, 1767225600}, + {1767225601, 1767225839, 1767225660, 1767225720}, + {-179, -1, -120, -120}, + } { + t.Run(fmt.Sprint(tc.lower), func(t *testing.T) { + query := strings.NewReplacer("1767225600", fmt.Sprint(tc.lower), "1767225720", fmt.Sprint(tc.upper)).Replace(mysqlBucketQuery) + analysis := a.Analyze(query, time.Time{}) + if analysis.Mode != sqlanalyzer.CacheModeDelta { + t.Fatalf("not delta cacheable: %+v", analysis) + } + extent := analysis.Plan.RequestExtent(time.Time{}) + if extent.Start.Unix() != tc.start || extent.End.Unix() != tc.end { + t.Fatalf("extent=%+v want=[%d,%d]", extent, tc.start, tc.end) + } + rendered, err := analysis.Plan.RenderExtent(extent) + if err != nil { + t.Fatal(err) + } + roundTrip := a.Analyze(rendered, time.Time{}) + if roundTrip.Mode != sqlanalyzer.CacheModeDelta || roundTrip.Plan.DropsPartialBuckets || roundTrip.Plan.CanonicalSQL != analysis.Plan.CanonicalSQL { + t.Fatalf("invalid normalized query: %s %+v", rendered, roundTrip) + } + }) + } +} diff --git a/pkg/backends/greptimedb/analyzer.go b/pkg/backends/greptimedb/analyzer.go new file mode 100644 index 000000000..915afcc38 --- /dev/null +++ b/pkg/backends/greptimedb/analyzer.go @@ -0,0 +1,96 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 greptimedb + +import ( + "errors" + "time" + + "github.com/trickstercache/trickster/v2/pkg/parsing/sqlanalyzer" + "github.com/trickstercache/trickster/v2/pkg/parsing/sqlanalyzer/cockroach" + "github.com/trickstercache/trickster/v2/pkg/parsing/sqlguard" + "github.com/trickstercache/trickster/v2/pkg/proxy/pgwire" +) + +var ( + errVolatile = errors.New("statement reads changing state") + compactUnits = map[string]time.Duration{ + "ns": time.Nanosecond, "us": time.Microsecond, "ms": time.Millisecond, + "s": time.Second, "m": time.Minute, "h": time.Hour, "d": 24 * time.Hour, "w": 7 * 24 * time.Hour, + } + guardedWords = map[string]sqlguard.WordClass{ + "random": sqlguard.Volatile, "rand": sqlguard.Volatile, "uuid": sqlguard.Volatile, + "uuid_v4": sqlguard.Volatile, "uuid_v7": sqlguard.Volatile, "gen_random_uuid": sqlguard.Volatile, + "jev": sqlguard.Volatile, "flush_flow": sqlguard.Volatile, "procedure_state": sqlguard.Volatile, + "now": sqlguard.Clock, "today": sqlguard.Clock, + "current_timestamp": sqlguard.BareClock, "current_date": sqlguard.BareClock, + "current_time": sqlguard.BareClock, "localtimestamp": sqlguard.BareClock, "localtime": sqlguard.BareClock, + "unnest": sqlguard.SetReturning, "generate_series": sqlguard.SetReturning, + "only": sqlguard.Unfaithful, + } + analyzer = &sessionAnalyzer{utc: newAnalyzer(true), zoned: newAnalyzer(false)} +) + +type sessionAnalyzer struct{ utc, zoned *dialectAnalyzer } + +var _ pgwire.SessionAnalyzer = (*sessionAnalyzer)(nil) + +func (a *sessionAnalyzer) Analyze(statement string, now time.Time) sqlanalyzer.Analysis { + return a.zoned.Analyze(statement, now) +} + +func (a *sessionAnalyzer) ForSession(view pgwire.SessionView) sqlanalyzer.DialectAnalyzer { + if view.UTC { + return a.utc + } + return a.zoned +} + +type dialectAnalyzer struct{ inner *cockroach.Analyzer } + +func newAnalyzer(utc bool) *dialectAnalyzer { + matchers := []cockroach.BucketMatcher{cockroach.DateBinMatcher, cockroach.CompactDateBinMatcher(compactUnits)} + if utc { + matchers = append(matchers, cockroach.DateTruncMatcher) + } + return &dialectAnalyzer{inner: cockroach.NewAnalyzer(cockroach.Options{ + BucketMatchers: matchers, ExprBucketMatchers: []cockroach.ExprBucketMatcher{cockroach.EpochFloorMatcher}, + RoundUnalignedTimeBounds: true, NakedIntIsInt4: true, + // DataFusion truncates finer bounds instead of rounding up as PostgreSQL does. + BoundPrecision: time.Nanosecond, RejectZonelessBounds: !utc, PostRender: postRender, + })} +} + +func (a *dialectAnalyzer) Analyze(statement string, now time.Time) sqlanalyzer.Analysis { + analysis := a.inner.Analyze(statement, now) + facts := sqlguard.Scan(statement, guardedWords) + if analysis.Mode == sqlanalyzer.CacheModeDelta && facts.Clock && !facts.Volatile { + facts.Clock = sqlguard.Scan(cockroach.MaskPlaceholders(analysis.Plan.CanonicalSQL), guardedWords).Clock + } + switch { + case facts.Volatile || facts.Clock: + return sqlanalyzer.Analysis{Mode: sqlanalyzer.CacheModeNone, Reason: sqlanalyzer.ReasonNondeterministic, Err: errVolatile} + case analysis.Mode != sqlanalyzer.CacheModeDelta: + return analysis + case facts.Unfaithful || facts.SetReturning: + return sqlanalyzer.ObjectAnalysis(sqlanalyzer.ReasonUnsupportedFormat, errUnrenderable) + case analysis.Plan.Step < time.Microsecond || analysis.Plan.Phase%time.Microsecond != 0: + // PostgreSQL text timestamps carry microseconds even for nanosecond columns. + return sqlanalyzer.ObjectAnalysis(sqlanalyzer.ReasonUnsupportedBucket, errUnrenderable) + } + return analysis +} diff --git a/pkg/backends/greptimedb/analyzer_test.go b/pkg/backends/greptimedb/analyzer_test.go new file mode 100644 index 000000000..639266290 --- /dev/null +++ b/pkg/backends/greptimedb/analyzer_test.go @@ -0,0 +1,159 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 greptimedb + +import ( + "strings" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/parsing/sqlanalyzer" + "github.com/trickstercache/trickster/v2/pkg/proxy/pgwire" + "github.com/trickstercache/trickster/v2/pkg/timeseries" +) + +const testRange = " FROM trips WHERE pickup_datetime >= '2026-01-01T00:00:00Z' AND pickup_datetime < '2026-02-01T00:00:00Z' GROUP BY 1 ORDER BY 1" + +func TestCacheAnalyzer(t *testing.T) { + a := Engine().Analyzer() + if a == nil { + t.Fatal("GreptimeDB must analyze SQL before enabling its native cache") + } + if sessions, ok := a.(pgwire.SessionAnalyzer); ok { + a = sessions.ForSession(pgwire.SessionView{UTC: true}) + } + for _, bucket := range []string{ + "date_bin(INTERVAL '5 minutes', pickup_datetime)", + "date_bin('5m', pickup_datetime)", + "date_trunc('week', pickup_datetime)", + "floor(extract(epoch FROM pickup_datetime)/300)*300", + "floor(date_part('epoch', pickup_datetime)/300)*300", + } { + t.Run(bucket, func(t *testing.T) { + analysis := a.Analyze("SELECT "+bucket+" AS time, count(*)"+testRange, time.Now()) + if analysis.Mode != sqlanalyzer.CacheModeDelta { + t.Fatalf("got %v/%v: %v", analysis.Mode, analysis.Reason, analysis.Err) + } + }) + } +} + +func TestAnalyzerFailsClosed(t *testing.T) { + for name, test := range map[string]struct { + sql string + mode sqlanalyzer.CacheMode + reason sqlanalyzer.AnalysisReason + }{ + "scalar": {"SELECT 1", sqlanalyzer.CacheModeObject, sqlanalyzer.ReasonUnsupportedFormat}, + "limited": {"SELECT * FROM trips LIMIT 10", sqlanalyzer.CacheModeObject, sqlanalyzer.ReasonUnsupportedLimit}, + "insert": {"INSERT INTO trips VALUES (1)", sqlanalyzer.CacheModeNone, sqlanalyzer.ReasonUnsupportedStatement}, + "ddl": {"CREATE TABLE a (b int)", sqlanalyzer.CacheModeNone, sqlanalyzer.ReasonUnsupportedStatement}, + "tql": {"TQL EVAL (0, 1, '1s') up", sqlanalyzer.CacheModeNone, sqlanalyzer.ReasonInvalidSQL}, + "admin": {"ADMIN FLUSH_TABLE('trips')", sqlanalyzer.CacheModeNone, sqlanalyzer.ReasonInvalidSQL}, + "range select": {"SELECT sum(fare_amount) RANGE '5m' FROM trips ALIGN '5m'", sqlanalyzer.CacheModeObject, sqlanalyzer.ReasonInvalidSQL}, + "volatile despite limit": {"SELECT random() FROM trips LIMIT 10", sqlanalyzer.CacheModeNone, sqlanalyzer.ReasonNondeterministic}, + "volatile in cte": {"WITH x AS (SELECT random()) SELECT * FROM x", sqlanalyzer.CacheModeNone, sqlanalyzer.ReasonNondeterministic}, + "volatile in unparsed dialect": {"SELECT random() RANGE '5m' FROM trips ALIGN '5m'", sqlanalyzer.CacheModeNone, sqlanalyzer.ReasonNondeterministic}, + "clock output": {"SELECT current_timestamp", sqlanalyzer.CacheModeNone, sqlanalyzer.ReasonNondeterministic}, + "side effect": {"SELECT flush_flow('f')", sqlanalyzer.CacheModeNone, sqlanalyzer.ReasonNondeterministic}, + "calendar bucket": {"SELECT date_trunc('month', pickup_datetime) AS time, count(*)" + testRange, sqlanalyzer.CacheModeObject, sqlanalyzer.ReasonUnsupportedBucket}, + "fractional compact width": {"SELECT date_bin('0.5s', pickup_datetime) AS time, count(*)" + testRange, sqlanalyzer.CacheModeObject, sqlanalyzer.ReasonUnsupportedBucket}, + "submicrosecond output": {"SELECT date_bin('5ns', pickup_datetime) AS time, count(*)" + testRange, sqlanalyzer.CacheModeObject, sqlanalyzer.ReasonUnsupportedBucket}, + } { + t.Run(name, func(t *testing.T) { + a := analyzer.ForSession(pgwire.SessionView{UTC: true}).Analyze(test.sql, time.Now()) + if a.Mode != test.mode || a.Reason != test.reason { + t.Fatalf("got %s/%s (%v), want %s/%s", a.Mode, a.Reason, a.Err, test.mode, test.reason) + } + }) + } +} + +func TestAnalyzerRender(t *testing.T) { + for _, bucket := range []string{"date_bin(INTERVAL '5 minutes', pickup_datetime)", "date_bin('5m', pickup_datetime)", "floor(extract(epoch FROM pickup_datetime)/300)*300"} { + sql := "SELECT " + bucket + " AS time, count(*)" + testRange + a := analyzer.ForSession(pgwire.SessionView{UTC: true}).Analyze(sql, time.Now()) + if a.Mode != sqlanalyzer.CacheModeDelta { + t.Fatalf("%s: %+v", sql, a) + } + extent := timeseries.Extent{Start: time.Date(2026, 1, 2, 0, 0, 0, 0, time.UTC), End: time.Date(2026, 1, 2, 1, 0, 0, 0, time.UTC)} + for range 2 { + rendered, err := a.Plan.RenderExtent(extent) + if err != nil { + t.Fatal(err) + } + if strings.Contains(rendered, "extract('epoch',") || !strings.Contains(rendered, "2026-01-02") { + t.Fatalf("bad rendering: %s", rendered) + } + b := analyzer.ForSession(pgwire.SessionView{UTC: true}).Analyze(rendered, time.Now()) + if b.Mode != sqlanalyzer.CacheModeDelta || b.Plan.Step != a.Plan.Step { + t.Fatalf("cannot analyze rendered SQL %s: %+v", rendered, b) + } + } + } + sql := "SELECT date_bin('5m', pickup_datetime) AS time, count(*)" + strings.ReplaceAll(testRange, "T00:00:00Z", " 00:00:00") + if a := analyzer.Analyze(sql, time.Now()); a.Mode == sqlanalyzer.CacheModeDelta { + t.Fatal("unknown session zone accepted naive bounds") + } + if a := analyzer.ForSession(pgwire.SessionView{UTC: true}).Analyze(sql, time.Now()); a.Mode != sqlanalyzer.CacheModeDelta { + t.Fatalf("known UTC rejected naive bounds: %+v", a) + } +} + +func TestPostRender(t *testing.T) { + for input, want := range map[string]string{ + "SELECT extract('epoch', ts)": "SELECT extract(epoch FROM ts)", + "SELECT bucket::STRING, 'bucket STRING', col::BYTES": "SELECT \"bucket\"::TEXT, 'bucket STRING', col::BYTEA", + "SELECT 'extract(''epoch'', ts)'": "SELECT 'extract(''epoch'', ts)'", + "SELECT extract(epoch FROM ts)": "SELECT extract(epoch FROM ts)", + "SELECT extract": "SELECT extract", + } { + got, err := postRender(input) + if err != nil || got != want { + t.Fatalf("%s: got %s, %v; want %s", input, got, err, want) + } + } + for _, input := range []string{"SELECT extract('', ts)", "SELECT extract('bad-field', ts)", "SELECT extract('a''b', ts)"} { + if _, err := postRender(input); err == nil { + t.Fatalf("accepted invalid extract field: %s", input) + } + } +} + +func TestAnalyzerGuardedBuckets(t *testing.T) { + for name, tc := range map[string]struct { + sql string + mode sqlanalyzer.CacheMode + reason sqlanalyzer.AnalysisReason + }{ + "clock resolved only in bounds": {"SELECT date_bin('5m', pickup_datetime) AS time, count(*) FROM trips WHERE pickup_datetime >= now() - INTERVAL '1 hour' AND pickup_datetime < now() GROUP BY 1", sqlanalyzer.CacheModeDelta, sqlanalyzer.ReasonDeltaCacheable}, + "unfaithful qualifier": {"SELECT date_bin('5m', pickup_datetime) AS time, count(*)" + strings.Replace(testRange, "FROM trips", "FROM ONLY trips", 1), sqlanalyzer.CacheModeObject, sqlanalyzer.ReasonUnsupportedFormat}, + "set returning": {"SELECT date_bin('5m', pickup_datetime) AS time, max(unnest(items))" + testRange, sqlanalyzer.CacheModeObject, sqlanalyzer.ReasonUnsupportedFormat}, + "submicrosecond origin": {"SELECT date_bin('5m', pickup_datetime, TIMESTAMP '1970-01-01T00:00:00.000000001Z') AS time, count(*)" + testRange, sqlanalyzer.CacheModeObject, sqlanalyzer.ReasonUnsupportedBucket}, + } { + t.Run(name, func(t *testing.T) { + a := analyzer.ForSession(pgwire.SessionView{UTC: true}).Analyze(tc.sql, time.Now()) + if a.Mode != tc.mode || a.Reason != tc.reason { + t.Fatalf("got %s/%s (%v), want %s/%s", a.Mode, a.Reason, a.Err, tc.mode, tc.reason) + } + }) + } + sql := "SELECT floor(extract(epoch FROM pickup_datetime)/300)*300 AS time, count(*)" + testRange + if a := analyzer.ForSession(pgwire.SessionView{}).Analyze(sql, time.Now()); a.Mode != sqlanalyzer.CacheModeDelta { + t.Fatalf("absolute bounds must remain cacheable without a UTC session: %+v", a) + } +} diff --git a/pkg/backends/greptimedb/cache_contract_test.go b/pkg/backends/greptimedb/cache_contract_test.go new file mode 100644 index 000000000..e2b60e447 --- /dev/null +++ b/pkg/backends/greptimedb/cache_contract_test.go @@ -0,0 +1,130 @@ +/* + * Copyright 2026 The Trickster Authors + * 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 greptimedb + +import ( + "errors" + "fmt" + "io" + "net/http" + "net/url" + "strings" + "sync/atomic" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/timeseries" +) + +func TestPrometheusCredentialIsolation(t *testing.T) { + for _, method := range []string{http.MethodGet, http.MethodPost} { + for _, endpoint := range []string{"query", "series", "labels", "label/job/values"} { + if method == http.MethodPost && endpoint == "label/job/values" { + continue + } + for _, policy := range []string{"", "private, max-age=60", "public, max-age=60"} { + t.Run(method+"/"+endpoint+"/"+policy, func(t *testing.T) { + var calls atomic.Int32 + shared := strings.HasPrefix(policy, "public") + h := newHTTPHarness(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + calls.Add(1) + user := r.Header.Get("Authorization") + if shared { + user = "public" + } + w.Header().Set("Content-Type", "application/json") + if policy != "" { + w.Header().Set("Cache-Control", policy) + } + fmt.Fprintf(w, `{"status":"success","data":[%q]}`, user) + })) + values := url.Values{"query": {"up"}, "match[]": {"up"}, "time": {"1704067200"}} + for _, user := range []string{"Basic dXNlcjph", "Basic dXNlcjpi", "Basic dXNlcjph", "Basic dXNlcjpi", ""} { + w := h.promQuery(t, method, endpoint, values, http.Header{"Authorization": {user}}) + want := user + if shared { + want = "public" + } + if body := fmt.Sprintf(`{"status":"success","data":[%q]}`, want); w.Code != 200 || w.Body.String() != body { + t.Fatalf("credential-specific response changed: code=%d body=%s want=%s", w.Code, w.Body.String(), body) + } + } + if shared && method == http.MethodGet && calls.Load() != 1 { + t.Fatalf("explicit origin sharing was lost: %d origin requests", calls.Load()) + } + }) + } + } + } +} + +func TestHTTPSQLMarshalFailurePreservesFallback(t *testing.T) { + for _, method := range []string{http.MethodGet, http.MethodPost} { + t.Run(method, func(t *testing.T) { + origin := &httpOrigin{} + h := newHTTPHarness(t, origin) + marshal := h.client.sqlModeler.WireMarshalWriter + h.client.sqlModeler.WireMarshalWriter = func(_ timeseries.Timeseries, _ *timeseries.RequestOptions, _ int, w io.Writer) error { + _, _ = io.WriteString(w, "partial invalid result") + return errors.New("injected marshal failure") + } + start := time.Now().UTC().Truncate(15 * time.Minute).Add(-3 * time.Hour) + statement := liveRange(start, start.Add(time.Hour)) + w := h.query(t, method, statement, nil, nil) + assertHTTPResult(t, w, "HTTPProxy", "proxy-only", 4) + if strings.Contains(w.Body.String(), "partial invalid result") || w.Header().Get("X-Greptime-Execution-Time") != "37" { + t.Fatal("failed pre-render leaked into the original response") + } + if got := origin.snapshot(); len(got) != 2 || got[1] != statement { + t.Fatalf("fallback did not replay original SQL: %v", got) + } + h.client.sqlModeler.WireMarshalWriter = marshal + assertHTTPResult(t, h.query(t, method, statement, nil, nil), "DeltaProxyCache", "kmiss", 4) + assertHTTPResult(t, h.query(t, method, statement, nil, nil), "DeltaProxyCache", "hit", 4) + }) + } +} + +func TestHTTPSQLSerializesOnce(t *testing.T) { + for _, method := range []string{http.MethodGet, http.MethodPost} { + t.Run(method, func(t *testing.T) { + h := newHTTPHarness(t, &httpOrigin{}) + var calls atomic.Int32 + marshal := h.client.sqlModeler.WireMarshalWriter + h.client.sqlModeler.WireMarshalWriter = func(ts timeseries.Timeseries, ro *timeseries.RequestOptions, status int, w io.Writer) error { + calls.Add(1) + return marshal(ts, ro, status, w) + } + start := time.Now().UTC().Truncate(15 * time.Minute).Add(-3 * time.Hour) + for _, tc := range []struct { + status string + end time.Time + rows int + header http.Header + }{ + {"kmiss", start.Add(time.Hour), 4, nil}, + {"hit", start.Add(time.Hour), 4, nil}, + {"phit", start.Add(75 * time.Minute), 5, nil}, + {"purge", start.Add(75 * time.Minute), 5, http.Header{"Cache-Control": {"no-cache"}}}, + } { + before := calls.Load() + w := h.query(t, method, liveRange(start, tc.end), nil, tc.header) + assertHTTPResult(t, w, "DeltaProxyCache", tc.status, tc.rows) + if got := calls.Load() - before; got != 1 { + t.Errorf("%s serialized %d times, want 1", tc.status, got) + } + } + }) + } +} diff --git a/pkg/backends/greptimedb/compatibility_test.go b/pkg/backends/greptimedb/compatibility_test.go new file mode 100644 index 000000000..d38c4e988 --- /dev/null +++ b/pkg/backends/greptimedb/compatibility_test.go @@ -0,0 +1,41 @@ +/* + * Copyright 2026 The Trickster Authors + * 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 greptimedb + +import ( + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/parsing/sqlanalyzer" + "github.com/trickstercache/trickster/v2/pkg/proxy/pgwire" + "github.com/trickstercache/trickster/v2/pkg/testutil/sqlcompat" +) + +const corpusPath = "testdata/compatibility/v1.json" + +func corpusAnalyze(zone, sql string) sqlanalyzer.Analysis { + return analyzer.ForSession(pgwire.SessionView{UTC: zone == "UTC"}).Analyze(sql, time.Time{}) +} + +func TestCompatibilityCorpus(t *testing.T) { + sqlcompat.Run(t, corpusPath, corpusAnalyze) +} + +func TestCompatibilityCorpusCoversGrafanaMacros(t *testing.T) { + sqlcompat.CheckGrafanaMacros(t, corpusPath) +} + +func BenchmarkCompatibilityCorpus(b *testing.B) { + sqlcompat.Benchmark(b, corpusPath, corpusAnalyze) +} diff --git a/pkg/backends/greptimedb/greptimedb.go b/pkg/backends/greptimedb/greptimedb.go new file mode 100644 index 000000000..9ba6be330 --- /dev/null +++ b/pkg/backends/greptimedb/greptimedb.go @@ -0,0 +1,87 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 greptimedb provides the GreptimeDB HTTP, PostgreSQL and MySQL backend. +package greptimedb + +import ( + "net/http" + + "github.com/trickstercache/trickster/v2/pkg/backends" + "github.com/trickstercache/trickster/v2/pkg/backends/greptimedb/model" + bo "github.com/trickstercache/trickster/v2/pkg/backends/options" + "github.com/trickstercache/trickster/v2/pkg/backends/prometheus" + "github.com/trickstercache/trickster/v2/pkg/backends/providers" + "github.com/trickstercache/trickster/v2/pkg/backends/providers/registry/types" + "github.com/trickstercache/trickster/v2/pkg/cache" + "github.com/trickstercache/trickster/v2/pkg/parsing/sqlanalyzer" + "github.com/trickstercache/trickster/v2/pkg/proxy/pgwire" + pgo "github.com/trickstercache/trickster/v2/pkg/proxy/pgwire/options" + "github.com/trickstercache/trickster/v2/pkg/timeseries" +) + +// DefaultPort is GreptimeDB's PostgreSQL wire-protocol port. +const DefaultPort = "4003" + +// Client serves one GreptimeDB backend over HTTP and native SQL protocols. +type Client struct { + *prometheus.Client + sqlModeler *timeseries.Modeler +} + +var ( + _ backends.Backend = (*Client)(nil) + _ backends.TimeseriesBackend = (*Client)(nil) + _ types.NewBackendClientFunc = NewClient +) + +// NewClient returns a GreptimeDB backend client. +func NewClient(name string, o *bo.Options, router http.Handler, + cache cache.Cache, _ backends.Backends, _ types.Lookup, +) (backends.Backend, error) { + c := &Client{sqlModeler: model.NewModeler()} + b, err := prometheus.NewClientWithHooks(name, o, router, cache, promHooks()) + c.Client = b + if err == nil { + c.RegisterHandlers(nil) + } + return c, err +} + +type engine struct{} + +var ( + _ pgwire.Engine = engine{} + _ pgwire.HTTPEngine = engine{} + _ pgwire.SessionDefaultsEngine = engine{} + _ pgwire.SessionSettingsEngine = engine{} +) + +// Engine returns the GreptimeDB PostgreSQL wire-protocol engine. +func Engine() pgwire.Engine { return engine{} } + +func (engine) Name() string { return providers.GreptimeDB } +func (engine) DefaultPort() string { return DefaultPort } +func (engine) Dialect() string { return providers.GreptimeDB } +func (engine) SupportsHTTP() bool { return true } +func (engine) Analyzer() sqlanalyzer.DialectAnalyzer { return analyzer } +func (engine) Defaults() pgwire.EngineDefaults { + return pgwire.EngineDefaults{UpstreamTLSMode: pgo.TLSModeDisable} +} + +func (engine) TimeAxis(oid uint32) (pgwire.TimeAxisKind, bool) { + return pgwire.StandardTimeAxis(oid) +} diff --git a/pkg/backends/greptimedb/greptimedb_test.go b/pkg/backends/greptimedb/greptimedb_test.go new file mode 100644 index 000000000..7dfdd80b1 --- /dev/null +++ b/pkg/backends/greptimedb/greptimedb_test.go @@ -0,0 +1,197 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 greptimedb + +import ( + "context" + "errors" + "net/http" + "slices" + "testing" + + mo "github.com/trickstercache/trickster/v2/pkg/backends/mysql/options" + bo "github.com/trickstercache/trickster/v2/pkg/backends/options" + "github.com/trickstercache/trickster/v2/pkg/backends/providers" + "github.com/trickstercache/trickster/v2/pkg/config/listener" + "github.com/trickstercache/trickster/v2/pkg/proxy/methods" + "github.com/trickstercache/trickster/v2/pkg/proxy/paths/matching" + "github.com/trickstercache/trickster/v2/pkg/proxy/pgwire" + pgo "github.com/trickstercache/trickster/v2/pkg/proxy/pgwire/options" + "github.com/trickstercache/trickster/v2/pkg/proxy/request" + tu "github.com/trickstercache/trickster/v2/pkg/testutil" +) + +func TestClientContract(t *testing.T) { + o := bo.New() + o.Provider, o.OriginURL = providers.GreptimeDB, "http://db.example:4000/prefix" + if err := o.Initialize("greptime"); err != nil { + t.Fatal(err) + } + backend, err := NewClient(o.Name, o, http.NotFoundHandler(), nil, nil, nil) + if err != nil { + t.Fatal(err) + } + if backend.Configuration() != o || backend.Name() != o.Name { + t.Fatal("lost backend identity") + } + paths := backend.DefaultPathConfigs(o) + if len(paths) != 7 || paths[0].Path != "/" || paths[0].HandlerName != "proxy" || + paths[0].MatchType != matching.PathMatchTypePrefix || !slices.Equal(paths[0].Methods, methods.AllHTTPMethods()) { + t.Fatalf("unexpected paths: %+v", paths) + } + paths[0].Path = "/changed" + if backend.DefaultPathConfigs(o)[0].Path != "/" { + t.Fatal("default paths share mutable state") + } + if handlers := backend.Handlers(); len(handlers) != 10 || handlers["proxy"] == nil || handlers["health"] == nil || handlers["query"] == nil || handlers["sql"] == nil { + t.Fatal("missing HTTP handlers") + } +} + +func TestEngineContract(t *testing.T) { + e := Engine() + if e.Name() != providers.GreptimeDB || e.Dialect() != providers.GreptimeDB || e.DefaultPort() != DefaultPort || + e.Analyzer() == nil || e.Defaults().UpstreamTLSMode != pgo.TLSModeDisable || !e.TimeSemantics().LosslessFloatText { + t.Fatal("unexpected engine contract") + } + if !e.(pgwire.HTTPEngine).SupportsHTTP() { + t.Fatal("engine must expose HTTP") + } + for _, oid := range []uint32{pgwire.OIDTimestamp, pgwire.OIDTimestampTZ, pgwire.OIDDate, pgwire.OIDInt8, 0} { + want, wantOK := pgwire.StandardTimeAxis(oid) + got, gotOK := e.TimeAxis(oid) + if got != want || gotOK != wantOK { + t.Fatalf("OID %d: incompatible time axis", oid) + } + } +} + +func TestSessionContract(t *testing.T) { + e := Engine() + probe := e.(pgwire.SessionDefaultsEngine).SessionDefaultsProbe() + if probe.SQL != "SHOW TIMEZONE; SHOW DateStyle; SHOW IntervalStyle" || + !slices.Equal(probe.Names, []string{"timezone", "datestyle", "intervalstyle"}) { + t.Fatalf("unexpected defaults probe: %+v", probe) + } + settings := e.(pgwire.SessionSettingsEngine).SessionSettings() + if !settings.LocalPersists || settings.Aliases["time_zone"] != "timezone" { + t.Fatal("lost GreptimeDB's SET LOCAL or time_zone semantics") + } + for _, name := range []string{"timezone", "datestyle", "intervalstyle", "bytea_output", "search_path"} { + if _, ok := settings.Tracked[name]; !ok { + t.Fatalf("untracked setting: %s", name) + } + } + for _, name := range []string{"application_name", "extra_float_digits", "standard_conforming_strings"} { + if _, ok := settings.Neutral[name]; !ok { + t.Fatalf("no-op setting changes cache state: %s", name) + } + } + if e.TimeSemantics().AssumedTimeZone != "" { + t.Fatal("the server timezone is configurable, not an engine guarantee") + } +} + +func TestProxyHandler(t *testing.T) { + client := &Client{} + origin, w, r, _, err := tu.NewTestInstance("", client.DefaultPathConfigs, + http.StatusOK, "{}", nil, providers.GreptimeDB, "/v1/sql", "error") + if err != nil { + t.Fatal(err) + } + defer origin.Close() + rsc := request.GetResources(r) + backend, err := NewClient("test", rsc.BackendOptions, nil, nil, nil, nil) + if err != nil { + t.Fatal(err) + } + client = backend.(*Client) + rsc.BackendClient, rsc.BackendOptions.HTTPClient = client, client.HTTPClient() + client.ProxyHandler(w, r) + if w.Code != http.StatusOK { + t.Fatalf("unexpected proxy status %d", w.Code) + } +} + +func TestHealthProtocolSelection(t *testing.T) { + for _, httpListener := range []bool{true, false} { + o := bo.New() + o.Provider, o.OriginURL = providers.GreptimeDB, "http://db.example:4000/prefix" + o.HasHTTPListener = httpListener + o.Postgres = pgo.New() + o.Postgres.UpstreamURL = "postgres://origin:secret@127.0.0.1:9/public" + if err := o.Initialize("greptime"); err != nil { + t.Fatal(err) + } + backend, err := NewClient(o.Name, o, nil, nil, nil, nil) + if err != nil { + t.Fatal(err) + } + c := backend.(*Client) + health := c.DefaultHealthCheckConfig() + probe := c.HealthCheckProbe() + if httpListener { + if probe != nil || health.Scheme != "http" || health.Host != "db.example:4000" || health.Path != "/prefix/health" { + t.Fatal("mixed backend must use HTTP health") + } + } else { + if probe == nil || health.Host != "" || health.Scheme != "" || health.Path != "" { + t.Fatal("native-only backend must use pgwire health even with an HTTP origin_url") + } + o.Postgres.UpstreamURL = "://bad" + if err := c.HealthCheckProbe()(context.Background()); !errors.Is(err, errProbeConfig) { + t.Fatalf("expected sanitized invalid probe: %v", err) + } + } + } + backend, err := NewClient("nil-options", nil, nil, nil, nil, nil) + if err != nil { + t.Fatal(err) + } + c := backend.(*Client) + if c.DefaultHealthCheckConfig() == nil || !errors.Is(c.HealthCheckProbe()(context.Background()), errProbeConfig) { + t.Fatal("nil options must fail the probe without panicking") + } +} + +func TestMySQLHealthProtocolSelection(t *testing.T) { + o := bo.New() + o.Provider, o.OriginURL = providers.GreptimeDB, "http://db.example:4000" + o.NativeListenerProtocols = []string{listener.ProtocolMySQL} + o.MySQL = mo.New() + o.MySQL.UpstreamURL = "mysql://reader:dev-password@127.0.0.1:9/public" + o.Postgres = pgo.New() + o.Postgres.UpstreamURL = "://unused-postgres" + if err := o.Initialize("greptime"); err != nil { + t.Fatal(err) + } + backend, err := NewClient(o.Name, o, nil, nil, nil, nil) + if err != nil { + t.Fatal(err) + } + c := backend.(*Client) + probe := c.HealthCheckProbe() + ctx, cancel := context.WithCancel(context.Background()) + cancel() + if probe == nil || errors.Is(probe(ctx), errProbeConfig) { + t.Fatal("MySQL-only backend selected PostgreSQL health options") + } + o.MySQL.UpstreamURL = "://invalid-mysql" + if err := c.HealthCheckProbe()(ctx); !errors.Is(err, errProbeConfig) { + t.Fatalf("invalid MySQL health config was not sanitized: %v", err) + } +} diff --git a/pkg/backends/greptimedb/handler_query_test.go b/pkg/backends/greptimedb/handler_query_test.go new file mode 100644 index 000000000..9f0f57550 --- /dev/null +++ b/pkg/backends/greptimedb/handler_query_test.go @@ -0,0 +1,376 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 greptimedb + +import ( + "encoding/json" + "fmt" + "io" + "net/http" + "net/http/httptest" + "net/url" + "regexp" + "strings" + "sync" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/proxy/headers" + "github.com/trickstercache/trickster/v2/pkg/proxy/request" + tu "github.com/trickstercache/trickster/v2/pkg/testutil" +) + +type httpOrigin struct { + mu sync.Mutex + statements []string + fault string + failStart time.Time +} + +var httpBound = regexp.MustCompile(`ts\s*(>=|<)\s*(?:TIMESTAMPTZ\s*)?'([^']+)'`) + +func (o *httpOrigin) ServeHTTP(w http.ResponseWriter, r *http.Request) { + _ = r.ParseForm() + statement := r.URL.Query().Get("sql") + if _, exists := r.URL.Query()["sql"]; !exists { + statement = r.PostForm.Get("sql") + } + o.mu.Lock() + o.statements = append(o.statements, statement) + fault := o.fault + failStart := o.failStart + o.mu.Unlock() + w.Header().Set("Content-Type", "application/json") + w.Header().Set("X-Greptime-Execution-Time", "37") + w.Header().Set("X-Greptime-Metrics", `{"cpu":5}`) + if !strings.Contains(strings.ToLower(statement), "date_bin") || fault == "unsupported" { + fmt.Fprint(w, `{"output":[{"records":{"schema":{"column_schemas":[{"name":"value","data_type":"String"}]},"rows":[["original"]],"total_rows":1}}],"execution_time_ms":37}`) + return + } + bounds := httpBound.FindAllStringSubmatch(statement, -1) + if len(bounds) != 2 { + http.Error(w, "missing SQL bounds: "+statement, 400) + return + } + start, e1 := time.Parse(time.RFC3339Nano, bounds[0][2]) + end, e2 := time.Parse(time.RFC3339Nano, bounds[1][2]) + if e1 != nil || e2 != nil { + http.Error(w, "bad SQL timestamps: "+statement, 400) + return + } + if !failStart.IsZero() && start.Equal(failStart) { + http.Error(w, "fixture gap failed", http.StatusServiceUnavailable) + return + } + rows := make([][]any, 0) + if fault != "empty" { + for current := start; current.Before(end); current = current.Add(15 * time.Minute) { + rows = append(rows, []any{current.UnixNano(), "a", current.Unix() / 900}) + } + } + kind := "Int64" + if fault == "schema" { + kind = "Float64" + } + schema := []map[string]string{{"name": "time", "data_type": "TimestampNanosecond"}, {"name": "host", "data_type": "String"}, {"name": "value", "data_type": kind}} + _ = json.NewEncoder(w).Encode(map[string]any{"output": []any{map[string]any{"records": map[string]any{ + "schema": map[string]any{"column_schemas": schema}, "rows": rows, "total_rows": len(rows), + }}}, "execution_time_ms": 37}) +} + +func (o *httpOrigin) snapshot() []string { + o.mu.Lock() + defer o.mu.Unlock() + return append([]string(nil), o.statements...) +} + +func (o *httpOrigin) setFault(fault string) { o.mu.Lock(); o.fault = fault; o.mu.Unlock() } + +type httpHarness struct { + client *Client + resources *request.Resources +} + +func newHTTPHarness(t *testing.T, origin http.Handler) *httpHarness { + t.Helper() + ts := httptest.NewServer(origin) + t.Cleanup(ts.Close) + placeholder, _, req, _, err := tu.NewTestInstance("", (&Client{}).DefaultPathConfigs, 200, "{}", nil, "greptimedb", "/v1/sql", "error") + if err != nil { + t.Fatal(err) + } + t.Cleanup(placeholder.Close) + initial := request.GetResources(req) + o := initial.BackendOptions + o.OriginURL = ts.URL + if err := o.Initialize("default"); err != nil { + t.Fatal(err) + } + b, err := NewClient("default", o, nil, initial.CacheClient, nil, nil) + if err != nil { + t.Fatal(err) + } + c := b.(*Client) + o.HTTPClient = c.HTTPClient() + pc := c.DefaultPathConfigs(o).Match("GET", "/v1/sql") + return &httpHarness{c, request.NewResources(o, pc, initial.CacheConfig, initial.CacheClient, c, initial.Tracer)} +} + +func (h *httpHarness) query(t *testing.T, method, statement string, extra url.Values, hdr http.Header) *httptest.ResponseRecorder { + t.Helper() + values := url.Values{"sql": {statement}, "db": {"public"}} + for k, v := range extra { + values[k] = v + } + path := "/v1/sql" + var body io.Reader + if method == "POST" { + body = strings.NewReader(values.Encode()) + } else { + path += "?" + values.Encode() + } + r := httptest.NewRequest(method, "http://trickster"+path, body) + r.Header = hdr.Clone() + if r.Header == nil { + r.Header = make(http.Header) + } + if method == "POST" { + r.Header.Set("Content-Type", "application/x-www-form-urlencoded") + } + res := h.resources + r = request.SetResources(r, request.NewResources(res.BackendOptions, res.PathConfig, res.CacheConfig, res.CacheClient, h.client, res.Tracer)) + w := httptest.NewRecorder() + h.client.QueryHandler(w, r) + return w +} + +func liveRange(start, end time.Time) string { + return fmt.Sprintf("SELECT date_bin('15m', ts) AS time, host, SUM(value) AS value FROM metrics WHERE ts >= '%s' AND ts < '%s' GROUP BY 1,2 ORDER BY time,host", start.Format(time.RFC3339Nano), end.Format(time.RFC3339Nano)) +} + +func assertHTTPResult(t *testing.T, w *httptest.ResponseRecorder, engine, status string, rows int) { + t.Helper() + gotEngine, gotStatus := headers.ParseResultEngineStatus(w.Header().Get(headers.NameTricksterResult)) + if w.Code != 200 || gotEngine != engine || gotStatus != status { + t.Fatalf("HTTP %d engine=%s status=%s body=%s", w.Code, gotEngine, gotStatus, w.Body.String()) + } + if rows < 0 { + return + } + var doc struct { + Output []struct { + Records struct { + Rows [][]any + Total int `json:"total_rows"` + Schema struct { + Columns []any `json:"column_schemas"` + } + } + } + Execution int `json:"execution_time_ms"` + } + if err := json.Unmarshal(w.Body.Bytes(), &doc); err != nil { + t.Fatal(err) + } + if len(doc.Output) != 1 || len(doc.Output[0].Records.Rows) != rows || doc.Output[0].Records.Total != rows || len(doc.Output[0].Records.Schema.Columns) != 3 { + t.Fatalf("bad reconstructed shape: %s", w.Body.String()) + } + if engine == "DeltaProxyCache" && (doc.Execution != 0 || w.Header().Get("X-Greptime-Execution-Time") != "0" || w.Header().Get("X-Greptime-Metrics") != "") { + t.Fatalf("stale execution metadata: %s %v", w.Body.String(), w.Header()) + } +} + +func TestHTTPSQLCacheFlow(t *testing.T) { + for _, method := range []string{"GET", "POST"} { + t.Run(method, func(t *testing.T) { + origin := &httpOrigin{} + h := newHTTPHarness(t, origin) + start := time.Now().UTC().Truncate(15 * time.Minute).Add(-3 * time.Hour) + first := liveRange(start, start.Add(time.Hour)) + assertHTTPResult(t, h.query(t, method, first, nil, nil), "DeltaProxyCache", "kmiss", 4) + wide := liveRange(start.Add(-15*time.Minute), start.Add(75*time.Minute)) + assertHTTPResult(t, h.query(t, method, wide, nil, nil), "DeltaProxyCache", "phit", 6) + assertHTTPResult(t, h.query(t, method, wide, nil, nil), "DeltaProxyCache", "hit", 6) + assertHTTPResult(t, h.query(t, method, first, nil, nil), "DeltaProxyCache", "hit", 4) + if queries := origin.snapshot(); len(queries) != 3 { + t.Fatalf("expected only initial and two gap queries: %v", queries) + } + }) + } +} + +func TestHTTPSQLCacheIdentity(t *testing.T) { + for _, statement := range []string{"SELECT 9 AS value", liveRange(time.Now().UTC().Truncate(15*time.Minute).Add(-3*time.Hour), time.Now().UTC().Truncate(15*time.Minute).Add(-2*time.Hour))} { + t.Run(statement, func(t *testing.T) { + origin := &httpOrigin{} + h := newHTTPHarness(t, origin) + for _, test := range []struct { + name string + extra url.Values + headers http.Header + }{ + {"default", nil, nil}, + {"database", url.Values{"db": {"other"}}, nil}, + {"database_header", nil, http.Header{"X-Greptime-Db-Name": {"other"}}}, + {"timezone", nil, http.Header{"X-Greptime-Timezone": {"+08:00"}}}, + {"authorization", nil, http.Header{"Authorization": {"Basic Zm9vOmJhcg=="}}}, + {"greptime_auth", nil, http.Header{"X-Greptime-Auth": {"Basic YmFyOmJheg=="}}}, + } { + t.Run(test.name, func(t *testing.T) { + before := len(origin.snapshot()) + for i := 0; i < 2; i++ { + w := h.query(t, "POST", statement, test.extra, test.headers) + _, status := headers.ParseResultEngineStatus(w.Header().Get(headers.NameTricksterResult)) + want := "kmiss" + if i > 0 { + want = "hit" + } + if w.Code != 200 || status != want { + t.Fatalf("identity %s response: %d %s %s", test.name, w.Code, status, w.Body.String()) + } + } + if len(origin.snapshot()) != before+1 { + t.Fatal("identity did not fetch exactly once") + } + }) + } + }) + } +} + +func TestHTTPSQLFallbackFlow(t *testing.T) { + start := time.Now().UTC().Truncate(15 * time.Minute).Add(-3 * time.Hour) + for _, test := range []struct { + name, statement, fault string + extra url.Values + cacheControl string + }{ + {"unsupported_response", liveRange(start, start.Add(time.Hour)), "unsupported", nil, ""}, + {"uncached_unsupported", liveRange(start, start.Add(time.Hour)), "unsupported", nil, "no-cache"}, + {"multiple", "SELECT 1; SELECT 2", "", nil, ""}, + {"write", "DELETE FROM metrics", "", nil, ""}, + {"volatile", "SELECT random()", "", nil, ""}, + } { + t.Run(test.name, func(t *testing.T) { + origin := &httpOrigin{fault: test.fault} + h := newHTTPHarness(t, origin) + for range 2 { + w := h.query(t, "POST", test.statement, test.extra, http.Header{"Cache-Control": {test.cacheControl}}) + engine, _ := headers.ParseResultEngineStatus(w.Header().Get(headers.NameTricksterResult)) + if w.Code != 200 || engine == "DeltaProxyCache" || engine == "ObjectProxyCache" || !strings.Contains(w.Body.String(), "original") { + t.Fatalf("not proxied: %s %s", engine, w.Body.String()) + } + } + queries := origin.snapshot() + if queries[len(queries)-1] != test.statement { + t.Fatalf("fallback did not restore original SQL: %v", queries) + } + }) + } +} + +func TestHTTPSQLEmptyAndSchemaChangeFlow(t *testing.T) { + start := time.Now().UTC().Truncate(15 * time.Minute).Add(-3 * time.Hour) + for _, fault := range []string{"empty", "schema", "unsupported"} { + t.Run(fault, func(t *testing.T) { + origin := &httpOrigin{} + if fault == "empty" { + origin.setFault(fault) + } + h := newHTTPHarness(t, origin) + statement := liveRange(start, start.Add(time.Hour)) + rows := 4 + if fault == "empty" { + rows = 0 + } + assertHTTPResult(t, h.query(t, "POST", statement, nil, nil), "DeltaProxyCache", "kmiss", rows) + origin.setFault(fault) + wide := liveRange(start.Add(-15*time.Minute), start.Add(75*time.Minute)) + w := h.query(t, "POST", wide, nil, nil) + if fault == "empty" { + assertHTTPResult(t, w, "DeltaProxyCache", "phit", 0) + assertHTTPResult(t, h.query(t, "POST", wide, nil, nil), "DeltaProxyCache", "hit", 0) + } else { + engine, _ := headers.ParseResultEngineStatus(w.Header().Get(headers.NameTricksterResult)) + queries := origin.snapshot() + if w.Code != 200 || engine == "DeltaProxyCache" || queries[len(queries)-1] != wide { + t.Fatalf("schema or gap failure not proxied: %s %s %v", engine, w.Body.String(), queries) + } + } + }) + } +} + +func TestHTTPSQLSingleFailedGap(t *testing.T) { + start := time.Now().UTC().Truncate(15 * time.Minute).Add(-3 * time.Hour) + for _, method := range []string{"GET", "POST"} { + t.Run(method, func(t *testing.T) { + origin := &httpOrigin{failStart: start.Add(time.Hour)} + h := newHTTPHarness(t, origin) + first := liveRange(start, start.Add(time.Hour)) + assertHTTPResult(t, h.query(t, method, first, nil, nil), "DeltaProxyCache", "kmiss", 4) + wide := liveRange(start.Add(-15*time.Minute), start.Add(75*time.Minute)) + for range 2 { + assertHTTPResult(t, h.query(t, method, wide, nil, nil), "HTTPProxy", "proxy-only", 6) + queries := origin.snapshot() + if queries[len(queries)-1] != wide { + t.Fatal("a partial failure did not restore the complete original query") + } + } + assertHTTPResult(t, h.query(t, method, first, nil, nil), "DeltaProxyCache", "hit", 4) + }) + } +} + +func TestHTTPSQLShardedFlow(t *testing.T) { + for _, method := range []string{"GET", "POST"} { + t.Run(method, func(t *testing.T) { + origin := &httpOrigin{} + h := newHTTPHarness(t, origin) + h.resources.BackendOptions.DoesShard = true + h.resources.BackendOptions.MaxShardSizePoints = 2 + start := time.Now().UTC().Truncate(15 * time.Minute).Add(-3 * time.Hour) + statement := liveRange(start, start.Add(90*time.Minute)) + assertHTTPResult(t, h.query(t, method, statement, nil, nil), "DeltaProxyCache", "kmiss", 6) + assertHTTPResult(t, h.query(t, method, statement, nil, nil), "DeltaProxyCache", "hit", 6) + if queries := origin.snapshot(); len(queries) != 3 { + t.Fatalf("expected three two-point shards: %v", queries) + } + }) + } +} + +func TestHTTPSQLAuthenticatedGETPolicy(t *testing.T) { + for _, shareable := range []bool{false, true} { + t.Run(fmt.Sprint(shareable), func(t *testing.T) { + origin := &httpOrigin{} + h := newHTTPHarness(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if shareable { + w.Header().Set("Cache-Control", "public, max-age=60") + } + origin.ServeHTTP(w, r) + })) + for i := range 2 { + status := "kmiss" + if shareable && i == 1 { + status = "hit" + } + assertHTTPResult(t, h.query(t, "GET", "SELECT 9", nil, http.Header{"Authorization": {"Basic Zm9vOmJhcg=="}}), "ObjectProxyCache", status, -1) + } + }) + } +} diff --git a/pkg/backends/greptimedb/health.go b/pkg/backends/greptimedb/health.go new file mode 100644 index 000000000..49cbb7571 --- /dev/null +++ b/pkg/backends/greptimedb/health.go @@ -0,0 +1,61 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 greptimedb + +import ( + "context" + "errors" + "slices" + + "github.com/trickstercache/trickster/v2/pkg/backends/healthcheck" + ho "github.com/trickstercache/trickster/v2/pkg/backends/healthcheck/options" + "github.com/trickstercache/trickster/v2/pkg/backends/mysql" + "github.com/trickstercache/trickster/v2/pkg/config/listener" + "github.com/trickstercache/trickster/v2/pkg/proxy/pgwire" +) + +var errProbeConfig = errors.New("greptimedb health probe configuration is invalid") + +// DefaultHealthCheckConfig selects HTTP health only for HTTP listener mappings. +func (c *Client) DefaultHealthCheckConfig() *ho.Options { + o := ho.New() + if options := c.Configuration(); options != nil && options.HasHTTPListener { + u := c.BaseUpstreamURL() + o.Scheme, o.Host, o.Path = u.Scheme, u.Host, u.Path+"/health" + } + return o +} + +// HealthCheckProbe uses a native login probe for native-only deployments. +func (c *Client) HealthCheckProbe() healthcheck.Probe { + if options := c.Configuration(); options != nil && options.HasHTTPListener { + return nil + } + if o := c.Configuration(); o != nil && slices.Contains(o.NativeListenerProtocols, listener.ProtocolMySQL) && + !slices.Contains(o.NativeListenerProtocols, listener.ProtocolPostgres) { + probe, err := mysql.HealthCheckProbeForEngine(o, MySQLEngine()) + if err != nil { + return func(context.Context) error { return errProbeConfig } + } + return probe + } + config, err := pgwire.ConfigFromOptions(c.Configuration(), Engine()) + if err != nil { + return func(context.Context) error { return errProbeConfig } + } + return config.Probe +} diff --git a/pkg/backends/greptimedb/http.go b/pkg/backends/greptimedb/http.go new file mode 100644 index 000000000..f410ba9dd --- /dev/null +++ b/pkg/backends/greptimedb/http.go @@ -0,0 +1,95 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 greptimedb + +import ( + "net/http" + "strings" + + "github.com/trickstercache/trickster/v2/pkg/backends/greptimedb/sql" + "github.com/trickstercache/trickster/v2/pkg/proxy/engines" + "github.com/trickstercache/trickster/v2/pkg/proxy/headers" + "github.com/trickstercache/trickster/v2/pkg/proxy/request" + "github.com/trickstercache/trickster/v2/pkg/proxy/urls" + "github.com/trickstercache/trickster/v2/pkg/timeseries" +) + +// QueryHandler serves GreptimeDB's HTTP SQL endpoint. +func (c *Client) QueryHandler(w http.ResponseWriter, r *http.Request) { + r.URL = urls.BuildUpstreamURL(r, c.BaseUpstreamURL()) + engines.DeltaProxyCacheRequest(sqlResponseWriter{w}, r, c.sqlModeler) +} + +// Rebuilt SQL results have no reusable origin execution metrics. Apply these +// headers at write time because DPC can serialize to a singleflight buffer. +type sqlResponseWriter struct{ http.ResponseWriter } + +func (w sqlResponseWriter) Unwrap() http.ResponseWriter { return w.ResponseWriter } + +func (w sqlResponseWriter) prepare() { + engine, _ := headers.ParseResultEngineStatus(w.Header().Get(headers.NameTricksterResult)) + if engine == "DeltaProxyCache" { + w.Header().Set("X-Greptime-Execution-Time", "0") + w.Header().Set("X-Greptime-Format", "greptimedb_v1") + w.Header().Del("X-Greptime-Metrics") + } +} + +func (w sqlResponseWriter) WriteHeader(code int) { + w.prepare() + w.ResponseWriter.WriteHeader(code) +} + +func (w sqlResponseWriter) Write(body []byte) (int, error) { + w.prepare() + return w.ResponseWriter.Write(body) +} + +func (c *Client) ParseTimeRangeQuery(r *http.Request) (*timeseries.TimeRangeQuery, + *timeseries.RequestOptions, bool, error, +) { + if isPromRange(r) { + return c.Client.ParseTimeRangeQuery(r) + } + a := analyzer.zoned + if r != nil { + timezone := r.Header.Get("X-Greptime-Timezone") + if res := request.GetResources(r); res != nil && res.PathConfig != nil { + // Request parameter overrides run after extent rendering. They must + // not replace the statement or its bounds after cache analysis. + if len(res.PathConfig.RequestParams) > 0 { + return nil, nil, false, timeseries.ErrUnknownFormat + } + effective := r.Clone(r.Context()) + headers.UpdateRequestHeaders(effective, res.PathConfig.RequestHeaders) + timezone = effective.Header.Get("X-Greptime-Timezone") + } + // HTTP's absent timezone is UTC, unlike pgwire's configurable default. + switch strings.ToUpper(timezone) { + case "", "UTC", "ETC/UTC", "+00:00", "-00:00": + a = analyzer.utc + } + } + return sql.ParseTimeRangeQuery(r, a) +} + +func (c *Client) SetExtent(r *http.Request, trq *timeseries.TimeRangeQuery, extent *timeseries.Extent) error { + if isPromRange(r) { + return c.Client.SetExtent(r, trq, extent) + } + return sql.SetExtent(r, trq, extent) +} diff --git a/pkg/backends/greptimedb/http_test.go b/pkg/backends/greptimedb/http_test.go new file mode 100644 index 000000000..e000dc172 --- /dev/null +++ b/pkg/backends/greptimedb/http_test.go @@ -0,0 +1,240 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 greptimedb + +import ( + "io" + "maps" + "net/http" + "net/http/httptest" + "net/url" + "reflect" + "strings" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/backends" + bo "github.com/trickstercache/trickster/v2/pkg/backends/options" + "github.com/trickstercache/trickster/v2/pkg/parsing/sqlanalyzer" + po "github.com/trickstercache/trickster/v2/pkg/proxy/paths/options" + "github.com/trickstercache/trickster/v2/pkg/proxy/request" + "github.com/trickstercache/trickster/v2/pkg/timeseries" +) + +const httpSQL = "SELECT date_bin('15m', ts) AS time, host, SUM(value) AS value FROM metrics WHERE ts >= '2024-01-01T00:00:00Z' AND ts < '2024-01-01T01:00:00Z' GROUP BY 1,2 ORDER BY time,host" + +func TestHTTPSQLProvider(t *testing.T) { + o := bo.New() + o.Provider, o.OriginURL = "greptimedb", "http://localhost:4000" + if err := o.Initialize("greptime"); err != nil { + t.Fatal(err) + } + b, err := NewClient(o.Name, o, nil, nil, nil, nil) + if err != nil { + t.Fatal(err) + } + c, ok := b.(backends.TimeseriesBackend) + if !ok { + t.Fatal("GreptimeDB HTTP SQL does not implement the timeseries cache contract") + } + r := httptest.NewRequest("GET", "/v1/sql?"+url.Values{"sql": {httpSQL}, "db": {"public"}}.Encode(), nil) + trq, ro, canOPC, err := c.ParseTimeRangeQuery(r) + if err != nil || trq == nil || ro == nil || !canOPC || trq.Step != 15*time.Minute { + t.Fatalf("HTTP SQL was not delta analyzed: query=%+v options=%+v object=%v err=%v", trq, ro, canOPC, err) + } + if c.Modeler() == nil || !ro.FastForwardDisable || o.FastForwardDisable { + t.Fatal("missing modeler or disabled SQL fast-forward guard") + } +} + +func TestHTTPSQLRequestModes(t *testing.T) { + for _, test := range []struct { + name, statement, params string + object, delta bool + }{ + {"delta", httpSQL, "", true, true}, + {"partial_lower", strings.Replace(httpSQL, "00:00:00Z", "00:00:01Z", 1), "", true, true}, + {"partial_upper", strings.Replace(httpSQL, "01:00:00Z", "00:59:59Z", 1), "", true, true}, + {"inclusive_upper", strings.Replace(httpSQL, "ts <", "ts <=", 1), "", true, true}, + {"complete_inclusive_upper", strings.Replace(strings.Replace(httpSQL, "ts <", "ts <=", 1), "01:00:00Z", "00:59:59.999999999Z", 1), "", true, true}, + {"count", "SELECT COUNT(*) FROM metrics", "", true, false}, + {"limit", httpSQL, "&limit=1", true, false}, + {"csv", httpSQL, "&format=csv", true, false}, + {"unknown_format", httpSQL, "&format=future", true, false}, + {"unknown_parameter", httpSQL, "&future=1", true, false}, + {"format_case", httpSQL, "&format=GREPTIMEDB_V1", true, true}, + {"delete", "DELETE FROM metrics", "", false, false}, + {"insert", "INSERT INTO metrics VALUES (1)", "", false, false}, + {"multiple", httpSQL + "; SELECT 1", "", false, false}, + {"volatile", "SELECT random()", "", false, false}, + {"clock", "SELECT now()", "", false, false}, + {"duplicate", httpSQL, "&sql=SELECT+1", false, false}, + {"bad_escape", httpSQL, "&db=%XY", false, false}, + } { + t.Run(test.name, func(t *testing.T) { + r := httptest.NewRequest("GET", "/v1/sql?sql="+url.QueryEscape(test.statement)+test.params, nil) + trq, _, object, err := (&Client{}).ParseTimeRangeQuery(r) + if object != test.object || (err == nil) != test.delta { + t.Fatalf("object=%v delta=%v query=%+v err=%v", object, err == nil, trq, err) + } + }) + } +} + +func TestHTTPSQLFormRewrite(t *testing.T) { + for _, inURL := range []bool{false, true} { + t.Run(map[bool]string{false: "form", true: "url"}[inURL], func(t *testing.T) { + form := url.Values{"sql": {httpSQL}, "db": {"form_db"}, "format": {"greptimedb_v1"}, "epoch": {"ns"}} + query := url.Values{"db": {"url_db"}} + if inURL { + query.Set("sql", httpSQL) + form.Set("sql", "SELECT 5") + } + body := form.Encode() + r := httptest.NewRequest("POST", "/v1/sql?"+query.Encode(), strings.NewReader(body)) + r.Header.Set("Content-Type", "application/x-www-form-urlencoded; charset=utf-8") + c := &Client{} + trq, _, _, err := c.ParseTimeRangeQuery(r) + if err != nil { + t.Fatal(err) + } + identity := maps.Clone(trq.CacheKeyElements) + if string(trq.OriginalBody) != body || !strings.Contains(identity["greptime.http"], "db=url_db") { + t.Fatal("body or effective database identity was lost") + } + extent := trq.Extent + extent.Start = extent.Start.Add(15 * time.Minute) + wantSQL, err := trq.ParsedQuery.(*sqlanalyzer.QueryPlan).RenderExtent(extent) + if err != nil { + t.Fatal(err) + } + if err := c.SetExtent(r, trq, &extent); err != nil { + t.Fatal(err) + } + raw, _ := io.ReadAll(r.Body) + gotForm, err := url.ParseQuery(string(raw)) + if err != nil { + t.Fatal(err) + } + if inURL { + query.Set("sql", wantSQL) + } else { + form.Set("sql", wantSQL) + } + if !reflect.DeepEqual(query, r.URL.Query()) || !reflect.DeepEqual(form, gotForm) || !maps.Equal(identity, trq.CacheKeyElements) { + t.Fatalf("extent rewrite changed unrelated fields: %v %v", r.URL.Query(), gotForm) + } + }) + } +} + +func TestHTTPSQLRejectedRequestsPreserveBody(t *testing.T) { + for _, test := range []struct{ name, method, query, body, contentType string }{ + {"empty_url", "POST", "sql=", "sql=SELECT+1", "application/x-www-form-urlencoded"}, + {"duplicate_form", "POST", "sql=SELECT+1", "sql=SELECT+2&sql=SELECT+3", "application/x-www-form-urlencoded"}, + {"malformed_form", "POST", "sql=SELECT+1", "db=%XX", "application/x-www-form-urlencoded"}, + {"json", "POST", "sql=SELECT+1", `{"sql":"SELECT 2"}`, "application/json"}, + {"missing_type", "POST", "sql=SELECT+1", "", ""}, + {"method", "PUT", "sql=SELECT+1", "sql=SELECT+2", "application/x-www-form-urlencoded"}, + } { + t.Run(test.name, func(t *testing.T) { + r := httptest.NewRequest(test.method, "/v1/sql?"+test.query, strings.NewReader(test.body)) + r.Header.Set("Content-Type", test.contentType) + _, _, can, err := (&Client{}).ParseTimeRangeQuery(r) + if err == nil || can { + t.Fatal("invalid request became cacheable") + } + body, _ := io.ReadAll(r.Body) + if string(body) != test.body { + t.Fatal("fallback request body changed") + } + }) + } + if _, _, can, err := (&Client{}).ParseTimeRangeQuery(nil); err == nil || can { + t.Fatal("nil request accepted") + } + if err := (&Client{}).SetExtent(nil, nil, nil); err == nil { + t.Fatal("nil rewrite accepted") + } +} + +func TestHTTPSQLTimezoneAndOverrides(t *testing.T) { + statement := strings.ReplaceAll(httpSQL, "date_bin('15m', ts)", "date_trunc('hour', ts)") + for _, test := range []struct { + name, timezone, override string + delta bool + }{ + {"default", "", "", true}, + {"utc", "UTC", "", true}, + {"offset", "+08:00", "", false}, + {"configured_offset", "UTC", "+08:00", false}, + {"configured_utc", "+08:00", "UTC", true}, + } { + t.Run(test.name, func(t *testing.T) { + r := httptest.NewRequest("GET", "/v1/sql?sql="+url.QueryEscape(statement), nil) + r.Header.Set("X-Greptime-Timezone", test.timezone) + pc := po.New() + if test.override != "" { + pc.RequestHeaders["X-Greptime-Timezone"] = test.override + } + r = request.SetResources(r, request.NewResources(bo.New(), pc, nil, nil, nil, nil)) + _, _, _, err := (&Client{}).ParseTimeRangeQuery(r) + if (err == nil) != test.delta { + t.Fatalf("delta=%v error=%v", test.delta, err) + } + if r.Header.Get("X-Greptime-Timezone") != test.timezone { + t.Fatal("input headers mutated") + } + pc.RequestParams = map[string]string{"sql": "DELETE FROM metrics"} + if _, _, can, err := (&Client{}).ParseTimeRangeQuery(r); err == nil || can { + t.Fatal("late query override was cached") + } + }) + } +} + +func TestHTTPSQLOpenEndedBackfill(t *testing.T) { + statement := strings.Replace(httpSQL, " AND ts < '2024-01-01T01:00:00Z'", "", 1) + r := httptest.NewRequest(http.MethodGet, "/v1/sql?sql="+url.QueryEscape(statement), nil) + trq, _, _, err := (&Client{}).ParseTimeRangeQuery(r) + if err != nil || trq.BackfillTolerance != 15*time.Minute { + t.Fatalf("query=%+v err=%v", trq, err) + } +} + +func BenchmarkHTTPSQLParse(b *testing.B) { + r := httptest.NewRequest("GET", "/v1/sql?sql="+url.QueryEscape(httpSQL), nil) + for b.Loop() { + if _, _, _, err := (&Client{}).ParseTimeRangeQuery(r); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkHTTPSQLRender(b *testing.B) { + r := httptest.NewRequest("GET", "/v1/sql?sql="+url.QueryEscape(httpSQL), nil) + trq, _, _, err := (&Client{}).ParseTimeRangeQuery(r) + if err != nil { + b.Fatal(err) + } + extent := timeseries.Extent{Start: trq.Extent.Start, End: trq.Extent.Start} + for b.Loop() { + if err := (&Client{}).SetExtent(r, trq, &extent); err != nil { + b.Fatal(err) + } + } +} diff --git a/pkg/backends/greptimedb/model/benchmark_test.go b/pkg/backends/greptimedb/model/benchmark_test.go new file mode 100644 index 000000000..28f73dee5 --- /dev/null +++ b/pkg/backends/greptimedb/model/benchmark_test.go @@ -0,0 +1,50 @@ +/* + * Copyright 2026 The Trickster Authors + * 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 model + +import ( + "bytes" + "fmt" + "strings" + "testing" + "time" +) + +func BenchmarkHTTPDecode(b *testing.B) { + const count = 10000 + for _, series := range []int{10, 1000, count} { + b.Run(fmt.Sprintf("series=%d", series), func(b *testing.B) { + var rows strings.Builder + rows.WriteByte('[') + for i := range count { + if i != 0 { + rows.WriteByte(',') + } + fmt.Fprintf(&rows, `[9007199254740993,%d,"host-%d"]`, int64(i)*int64(time.Second), i%series) + } + rows.WriteByte(']') + body := []byte(envelope(rows.String(), count)) + trq := query() + trq.Extent.End = time.Unix(count, 0) + b.SetBytes(int64(len(body))) + b.ReportAllocs() + for b.Loop() { + ts, err := UnmarshalTimeseriesReader(bytes.NewReader(body), trq) + if err != nil || ts.ValueCount() != count { + b.Fatalf("decode failed: %v", err) + } + } + }) + } +} diff --git a/pkg/backends/greptimedb/model/marshal.go b/pkg/backends/greptimedb/model/marshal.go new file mode 100644 index 000000000..c23b808ff --- /dev/null +++ b/pkg/backends/greptimedb/model/marshal.go @@ -0,0 +1,198 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 model + +import ( + "bytes" + "cmp" + "encoding/json" + "io" + "math/big" + "net/http" + "slices" + "strings" + + "github.com/trickstercache/trickster/v2/pkg/timeseries" +) + +func MarshalTimeseries(ts timeseries.Timeseries, options *timeseries.RequestOptions, status int) ([]byte, error) { + var buf bytes.Buffer + if err := MarshalTimeseriesWriter(ts, options, status, &buf); err != nil { + return nil, err + } + return buf.Bytes(), nil +} + +func MarshalTimeseriesWriter(ts timeseries.Timeseries, _ *timeseries.RequestOptions, _ int, w io.Writer) error { + d, ok := ts.(*dataSet) + if !ok || d == nil || d.DataSet == nil || d.invalid || len(d.fields) == 0 || w == nil { + return timeseries.ErrInvalidBody + } + r := &records{Schema: schema{Columns: make([]column, len(d.fields))}, Rows: make([][]any, 0)} + for i, field := range d.fields { + r.Schema.Columns[i] = column{Name: field.Name, Type: field.SDataType} + } + for _, result := range d.Results { + if result == nil { + continue + } + for _, series := range result.SeriesList { + if series == nil { + continue + } + template := make([]any, len(d.fields)) + for i, field := range d.fields { + if field.Role != timeseries.RoleTag { + continue + } + encoded, ok := series.Header.Tags[field.Name] + if !ok { + return timeseries.ErrInvalidBody + } + decoder := json.NewDecoder(strings.NewReader(encoded)) + decoder.UseNumber() + var value any + if err := decoder.Decode(&value); err != nil { + return err + } + v, err := decodeValue(value, field.SDataType) + if err != nil { + return err + } + template[i] = v + } + for _, point := range series.Points { + row := slices.Clone(template) + vi := 0 + for i, field := range d.fields { + switch field.Role { + case timeseries.RoleTimestamp: + v, err := epochValue(int64(point.Epoch), field) + if err != nil { + return err + } + row[i] = v + case timeseries.RoleValue: + if vi >= len(point.Values) { + return timeseries.ErrInvalidBody + } + row[i] = point.Values[vi] + vi++ + } + } + if vi != len(point.Values) { + return timeseries.ErrInvalidBody + } + r.Rows = append(r.Rows, row) + } + } + } + if d.TimeRangeQuery != nil { + sortRows(r.Rows, d.fields, d.TimeRangeQuery.Ordering) + } + total, elapsed := uint64(len(r.Rows)), uint64(0) + r.Total = &total + if hw, ok := w.(http.ResponseWriter); ok { + hw.Header().Set("Content-Type", "application/json") + hw.Header().Set("X-Greptime-Format", "greptimedb_v1") + hw.Header().Set("X-Greptime-Execution-Time", "0") + hw.Header().Del("X-Greptime-Metrics") + } + return json.NewEncoder(w).Encode(response{Output: []output{{Records: r}}, ExecutionTime: &elapsed}) +} + +func sortRows(rows [][]any, fields timeseries.FieldDefinitions, ordering []timeseries.OrderTerm) { + type orderColumn struct { + timeseries.OrderTerm + index int + } + columns := make([]orderColumn, 0, len(ordering)) + numbers := make(map[json.Number]*big.Rat) + for _, term := range ordering { + index := slices.IndexFunc(fields, func(field timeseries.FieldDefinition) bool { return field.Name == term.Column }) + if index < 0 { + continue + } + columns = append(columns, orderColumn{term, index}) + for _, row := range rows { + if n, ok := row[index].(json.Number); ok && numbers[n] == nil { + numbers[n], _ = new(big.Rat).SetString(string(n)) + } + } + } + slices.SortStableFunc(rows, func(a, b []any) int { + for _, term := range columns { + av, bv := a[term.index], b[term.index] + if av == nil || bv == nil { + if av == nil && bv == nil { + continue + } + if (av == nil) == term.NullsFirst { + return -1 + } + return 1 + } + comparison := compareValue(av, bv, numbers) + if comparison != 0 { + if term.Descending { + return -comparison + } + return comparison + } + } + return 0 + }) +} + +func compareValue(a, b any, numbers map[json.Number]*big.Rat) int { + switch av := a.(type) { + case int64: + if bv, ok := b.(int64); ok { + return cmp.Compare(av, bv) + } + case uint64: + if bv, ok := b.(uint64); ok { + return cmp.Compare(av, bv) + } + case float64: + if bv, ok := b.(float64); ok { + return cmp.Compare(av, bv) + } + case string: + if bv, ok := b.(string); ok { + return cmp.Compare(av, bv) + } + case bool: + if bv, ok := b.(bool); ok { + if av == bv { + return 0 + } + if av { + return 1 + } + return -1 + } + case json.Number: + if bv, ok := b.(json.Number); ok { + af, bf := numbers[av], numbers[bv] + if af != nil && bf != nil { + return af.Cmp(bf) + } + } + } + return 0 +} diff --git a/pkg/backends/greptimedb/model/model.go b/pkg/backends/greptimedb/model/model.go new file mode 100644 index 000000000..b46c34fef --- /dev/null +++ b/pkg/backends/greptimedb/model/model.go @@ -0,0 +1,113 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 model preserves GreptimeDB's typed HTTP SQL result envelope. +package model + +import ( + "bytes" + "slices" + "time" + + "github.com/trickstercache/trickster/v2/pkg/timeseries" + "github.com/trickstercache/trickster/v2/pkg/timeseries/dataset" +) + +// Schema lives outside the series so empty and cropped results retain it. +// Data and extent operations still use the common DataSet implementation. +type dataSet struct { + *dataset.DataSet + fields timeseries.FieldDefinitions + invalid bool +} + +var _ timeseries.Timeseries = (*dataSet)(nil) + +func NewModeler() *timeseries.Modeler { + return timeseries.NewModeler(UnmarshalTimeseries, UnmarshalTimeseriesReader, + MarshalTimeseries, MarshalTimeseriesWriter, unmarshalCache, marshalCache) +} + +func (d *dataSet) Clone() timeseries.Timeseries { + return &dataSet{DataSet: d.DataSet.Clone().(*dataset.DataSet), fields: d.fields.Clone(), invalid: d.invalid} +} + +func (d *dataSet) CroppedClone(extent timeseries.Extent) timeseries.Timeseries { + clone := d.Clone().(*dataSet) + clone.CropToRange(extent) + return clone +} + +func (d *dataSet) Merge(sortPoints bool, inputs ...timeseries.Timeseries) { + parts := make([]timeseries.Timeseries, 0, len(inputs)) + for _, input := range inputs { + part, ok := input.(*dataSet) + if !ok || part == nil || part.invalid || !slices.Equal(d.fields, part.fields) { + d.invalid = true + return + } + parts = append(parts, part.DataSet) + } + d.DataSet.Merge(sortPoints, parts...) +} + +func (d *dataSet) Size() int64 { + size := d.DataSet.Size() + for _, field := range d.fields { + size += int64(field.Size()) + } + return size +} + +var cacheVersion = []byte{'G', 'S', 'Q', 'L', 1} + +func marshalCache(ts timeseries.Timeseries, options *timeseries.RequestOptions, status int) ([]byte, error) { + d, ok := ts.(*dataSet) + if !ok || d == nil || d.invalid || d.DataSet == nil { + return nil, timeseries.ErrUnknownFormat + } + fields, err := d.fields.MarshalMsg(slices.Clone(cacheVersion)) + if err != nil { + return nil, err + } + body, err := dataset.MarshalDataSet(d.DataSet, options, status) + if err != nil { + return nil, err + } + return append(fields, body...), nil +} + +func unmarshalCache(body []byte, trq *timeseries.TimeRangeQuery) (timeseries.Timeseries, error) { + if !bytes.HasPrefix(body, cacheVersion) { + return nil, timeseries.ErrUnknownFormat + } + d := &dataSet{DataSet: &dataset.DataSet{}} + rest, err := d.fields.UnmarshalMsg(body[len(cacheVersion):]) + if err != nil || len(d.fields) == 0 { + return nil, timeseries.ErrInvalidBody + } + rest, err = d.DataSet.UnmarshalMsg(rest) + if err != nil || len(rest) != 0 { + return nil, timeseries.ErrInvalidBody + } + if trq != nil { + d.TimeRangeQuery = trq + } else if d.TimeRangeQuery != nil { + d.TimeRangeQuery.Step = time.Duration(d.TimeRangeQuery.StepNS) + d.TimeRangeQuery.PolicyStep = time.Duration(d.TimeRangeQuery.PolicyStepNS) + } + return d, nil +} diff --git a/pkg/backends/greptimedb/model/model_test.go b/pkg/backends/greptimedb/model/model_test.go new file mode 100644 index 000000000..95901132e --- /dev/null +++ b/pkg/backends/greptimedb/model/model_test.go @@ -0,0 +1,479 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 model + +import ( + "bytes" + "encoding/json" + "errors" + "fmt" + "io" + "net/http/httptest" + "reflect" + "strings" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/parsing/sqlanalyzer" + "github.com/trickstercache/trickster/v2/pkg/timeseries" + "github.com/trickstercache/trickster/v2/pkg/timeseries/dataset" +) + +func query() *timeseries.TimeRangeQuery { + plan := &sqlanalyzer.QueryPlan{ + CanonicalSQL: "fixture", OutputColumn: "time", GroupColumns: []string{"host"}, + ValueColumns: []string{"value"}, Step: time.Second, OutputUnit: timeseries.DateTimeRFC3339Nano, + Ordering: []timeseries.OrderTerm{{Column: "time", Descending: true}, {Column: "host", NullsFirst: true}}, + } + trq := sqlanalyzer.NewTimeRangeQuery("fixture") + plan.ApplyToQuery(trq) + trq.Extent = timeseries.Extent{Start: time.Unix(0, 0), End: time.Unix(10, 0)} + return trq +} + +func envelope(rows string, total int) string { + return fmt.Sprintf(`{"output":[{"records":{"schema":{"column_schemas":[{"name":"value","data_type":"UInt64"},{"name":"time","data_type":"TimestampNanosecond"},{"name":"host","data_type":"String"}]},"rows":%s,"total_rows":%d}}],"execution_time_ms":5}`, rows, total) +} + +func decoded(t testing.TB, body []byte) response { + t.Helper() + var out response + d := json.NewDecoder(bytes.NewReader(body)) + d.UseNumber() + if err := d.Decode(&out); err != nil { + t.Fatal(err) + } + return out +} + +func TestModelCacheRoundTrip(t *testing.T) { + input := envelope(`[[18446744073709551615,1000000000,"a"],[9007199254740993,2000000000,"null"],[null,2000000000,null],[7,2000000000,""]]`, 4) + trq := query() + ts, err := UnmarshalTimeseries([]byte(input), trq) + if err != nil { + t.Fatal(err) + } + if ts.SeriesCount() != 4 { + t.Fatalf("groups collapsed: %d", ts.SeriesCount()) + } + for _, cacheRoundTrip := range []bool{false, true} { + t.Run(fmt.Sprint(cacheRoundTrip), func(t *testing.T) { + candidate := ts.Clone() + if cacheRoundTrip { + buf, err := marshalCache(candidate, nil, 200) + if err != nil { + t.Fatal(err) + } + candidate, err = unmarshalCache(buf, trq) + if err != nil { + t.Fatal(err) + } + } + body, err := MarshalTimeseries(candidate, nil, 200) + if err != nil { + t.Fatal(err) + } + out := decoded(t, body) + want := decoded(t, []byte(envelope(`[[null,2000000000,null],[7,2000000000,""],[9007199254740993,2000000000,"null"],[18446744073709551615,1000000000,"a"]]`, 4))) + if !reflect.DeepEqual(out.Output, want.Output) || out.ExecutionTime == nil || *out.ExecutionTime != 0 { + t.Fatalf("fidelity lost: %s", body) + } + }) + } +} + +func TestSchemaSuppliesUnnamedValueColumns(t *testing.T) { + trq := query() + trq.ParsedQuery.(*sqlanalyzer.QueryPlan).ValueColumns = nil + if _, err := UnmarshalTimeseries([]byte(envelope(`[[7,1000000000,"a"]]`, 1)), trq); err != nil { + t.Fatal(err) + } + for _, bad := range []string{ + strings.Replace(envelope(`[]`, 0), `"name":"host"`, `"name":"other"`, 1), + strings.Replace(envelope(`[]`, 0), `"name":"value"`, `"name":"host"`, 1), + } { + if _, err := UnmarshalTimeseries([]byte(bad), trq); err == nil { + t.Fatal("missing or duplicate grouping column accepted") + } + } +} + +func TestModelEmptyAndCroppedSchema(t *testing.T) { + for _, rows := range []string{`[]`, `[[7,1000000000,"a"]]`} { + t.Run(rows, func(t *testing.T) { + total := 0 + if rows != "[]" { + total = 1 + } + trq := query() + ts, err := UnmarshalTimeseries([]byte(envelope(rows, total)), trq) + if err != nil { + t.Fatal(err) + } + cropped := ts.CroppedClone(timeseries.Extent{Start: time.Unix(3, 0), End: time.Unix(4, 0)}) + buf, err := marshalCache(cropped, nil, 200) + if err != nil { + t.Fatal(err) + } + cropped, err = unmarshalCache(buf, trq) + if err != nil { + t.Fatal(err) + } + body, err := MarshalTimeseries(cropped, nil, 200) + if err != nil { + t.Fatal(err) + } + got := decoded(t, body).Output[0].Records + if len(got.Schema.Columns) != 3 || got.Rows == nil || len(got.Rows) != 0 || *got.Total != 0 { + t.Fatalf("empty schema lost: %s", body) + } + if ts.ValueCount() != int64(total) { + t.Fatal("cropped clone mutated source") + } + }) + } +} + +func TestModelMerge(t *testing.T) { + trq := query() + a, err := UnmarshalTimeseries([]byte(envelope(`[]`, 0)), trq) + if err != nil { + t.Fatal(err) + } + b, err := UnmarshalTimeseries([]byte(envelope(`[[2,2000000000,"a"],[1,1000000000,"a"]]`, 2)), trq) + if err != nil { + t.Fatal(err) + } + a.Merge(true, b) + if a.ValueCount() != 2 { + t.Fatal("empty extent could not merge data") + } + bad := strings.Replace(envelope(`[[3,3000000000,"a"]]`, 1), "UInt64", "Int64", 1) + c, err := UnmarshalTimeseries([]byte(bad), trq) + if err != nil { + t.Fatal(err) + } + a.Merge(true, c) + if _, err := MarshalTimeseries(a, nil, 200); err == nil { + t.Fatal("incompatible schemas silently merged") + } + if _, err := marshalCache(a, nil, 200); err == nil { + t.Fatal("incompatible schema stored") + } +} + +func TestModelMalformedResponses(t *testing.T) { + base := envelope(`[[7,1000000000,"a"]]`, 1) + for name, input := range map[string]string{ + "empty": "", "trailing": base + "{}", "error": `{"code":1004,"error":"failed","execution_time_ms":0}`, + "affected_rows": `{"output":[{"affectedrows":1}],"execution_time_ms":0}`, + "multi": strings.Replace(base, `],"execution_time_ms"`, ` ,{"records":{}}],"execution_time_ms"`, 1), + "extra_field": strings.Replace(base, `"total_rows":1`, `"total_rows":1,"extra":true`, 1), + "metrics": strings.Replace(base, `"total_rows":1`, `"total_rows":1,"metrics":{"elapsed":1}`, 1), + "truncated": strings.Replace(base, `"total_rows":1`, `"total_rows":2`, 1), + "missing_total": strings.Replace(base, `,"total_rows":1`, "", 1), + "null_rows": envelope("null", 0), + "short_row": envelope(`[[7,1000000000]]`, 1), + "long_row": envelope(`[[7,1000000000,"a",1]]`, 1), + "duplicate_columns": strings.Replace(base, `"name":"host"`, `"name":"value"`, 1), + "wrong_column": strings.Replace(base, `"name":"host"`, `"name":"other"`, 1), + "unknown_type": strings.Replace(base, "UInt64", "Decimal128(38, 2)", 1), + "null_time": envelope(`[[7,null,"a"]]`, 1), + "fraction_time": envelope(`[[7,1.5,"a"]]`, 1), + "off_grid": envelope(`[[7,1000000001,"a"]]`, 1), + "overflow_time": envelope(`[[7,9223372036854775808,"a"]]`, 1), + "overflow_value": envelope(`[[18446744073709551616,1000000000,"a"]]`, 1), + "negative_unsigned": envelope(`[[-1,1000000000,"a"]]`, 1), + "wrong_tag_type": envelope(`[[7,1000000000,1]]`, 1), + "duplicate_point": envelope(`[[7,1000000000,"a"],[8,1000000000,"a"]]`, 2), + } { + t.Run(name, func(t *testing.T) { + if _, err := UnmarshalTimeseries([]byte(input), query()); err == nil { + t.Fatal("invalid response accepted") + } + }) + } +} + +func TestTypedValues(t *testing.T) { + for _, test := range []struct { + typ, value string + valid bool + }{ + {"Int8", "-128", true}, + {"Int8", "128", false}, + {"UInt8", "255", true}, + {"UInt8", "256", false}, + {"Int16", "-32768", true}, + {"Int32", "2147483647", true}, + {"Int64", "-9223372036854775808", true}, + {"UInt16", "65535", true}, + {"UInt32", "4294967295", true}, + {"UInt64", "18446744073709551615", true}, + {"Float32", "0.1", true}, + {"Float32", "3.5e38", false}, + {"Float64", "1.25e-20", true}, + {"Float64", "1e1000", false}, + {"String", `"a\u0000b"`, true}, + {"Boolean", "true", true}, + {"Boolean", "1", false}, + {"Null", "null", true}, + {"Null", "1", false}, + {"TimestampSecond", "-1", true}, + {"Date", "12345", true}, + } { + t.Run(test.typ+test.value, func(t *testing.T) { + d := json.NewDecoder(strings.NewReader(test.value)) + d.UseNumber() + var input any + if err := d.Decode(&input); err != nil { + t.Fatal(err) + } + v, err := decodeValue(input, test.typ) + if (err == nil) != test.valid { + t.Fatalf("value=%v err=%v", v, err) + } + }) + } +} + +func TestModelRejectsValuelessSeries(t *testing.T) { + trq := query() + trq.ParsedQuery.(*sqlanalyzer.QueryPlan).ValueColumns = nil + input := `{"output":[{"records":{"schema":{"column_schemas":[{"name":"time","data_type":"TimestampNanosecond"},{"name":"host","data_type":"String"}]},"rows":[[1000000000,"a"]],"total_rows":1}}],"execution_time_ms":0}` + if _, err := UnmarshalTimeseries([]byte(input), trq); err == nil { + t.Fatal("accepted a result whose rows would be discarded during cropping") + } +} + +func TestModelAmbiguousFloatOrdering(t *testing.T) { + input := strings.Replace(envelope(`[[null,1000000000,"a"],[1,2000000000,"a"]]`, 2), "UInt64", "Float64", 1) + trq := query() + if _, err := UnmarshalTimeseries([]byte(input), trq); err != nil { + t.Fatal("nullable float not used for sorting should be preserved", err) + } + plan := trq.ParsedQuery.(*sqlanalyzer.QueryPlan) + plan.Ordering = append(plan.Ordering, timeseries.OrderTerm{Column: "value"}) + if _, err := UnmarshalTimeseries([]byte(input), trq); err == nil { + t.Fatal("cannot distinguish NULL from non-finite floats when rebuilding sort order") + } +} + +func TestModelRequiresOrderingColumns(t *testing.T) { + trq := query() + plan := trq.ParsedQuery.(*sqlanalyzer.QueryPlan) + plan.ValueColumns = nil + plan.Ordering = []timeseries.OrderTerm{{Column: "missing_alias"}} + if _, err := UnmarshalTimeseries([]byte(envelope(`[[7,1000000000,"a"]]`, 1)), trq); err == nil { + t.Fatal("accepted a schema that cannot preserve the requested ordering") + } +} + +func TestExactNumericSort(t *testing.T) { + rows := [][]any{{json.Number("9007199254740993")}, {json.Number("9.007199254740992e15")}, {nil}} + fields := timeseries.FieldDefinitions{{Name: "time"}} + sortRows(rows, fields, []timeseries.OrderTerm{{Column: "time", NullsFirst: true}}) + if rows[0][0] != nil || rows[1][0] != json.Number("9.007199254740992e15") || rows[2][0] != json.Number("9007199254740993") { + t.Fatal("sort rounded distinct numbers", rows) + } +} + +func TestModelValueOrdering(t *testing.T) { + for _, test := range []struct { + typ, low, high string + }{ + {"Int64", "-9223372036854775808", "9223372036854775807"}, + {"UInt64", "9007199254740992", "9007199254740993"}, + {"Float64", "-1.25", "0.5"}, + {"String", `"a"`, `"b"`}, + {"Boolean", "false", "true"}, + } { + for _, descending := range []bool{false, true} { + for _, nullsFirst := range []bool{false, true} { + t.Run(fmt.Sprintf("%s/desc=%t/nullsFirst=%t", test.typ, descending, nullsFirst), func(t *testing.T) { + rows := fmt.Sprintf(`[[%s,1000000000,"a"],[%s,2000000000,"a"]`, test.high, test.low) + total := 2 + if test.typ != "Float64" { + rows += `,[null,3000000000,"a"]` + total++ + } + input := strings.Replace(envelope(rows+"]", total), "UInt64", test.typ, 1) + trq := query() + plan := trq.ParsedQuery.(*sqlanalyzer.QueryPlan) + plan.Ordering = []timeseries.OrderTerm{{Column: "value", Descending: descending, NullsFirst: nullsFirst}} + plan.ApplyToQuery(trq) + ts, err := UnmarshalTimeseries([]byte(input), trq) + if err != nil { + t.Fatal(err) + } + want := []json.Number{"2000000000", "1000000000"} + if descending { + want[0], want[1] = want[1], want[0] + } + if total == 3 { + if nullsFirst { + want = append([]json.Number{"3000000000"}, want...) + } else { + want = append(want, "3000000000") + } + } + for range 2 { + wire, err := MarshalTimeseries(ts, nil, 200) + if err != nil { + t.Fatal(err) + } + got := decoded(t, wire).Output[0].Records.Rows + if len(got) != len(want) { + t.Fatalf("row count changed: %s", wire) + } + for i, row := range got { + if row[1] != want[i] { + t.Fatalf("row order changed: %s", wire) + } + } + cached, err := marshalCache(ts, nil, 200) + if err != nil { + t.Fatal(err) + } + ts, err = unmarshalCache(cached, trq) + if err != nil { + t.Fatal(err) + } + } + }) + } + } + } +} + +func TestTimestampUnitsAndPrecision(t *testing.T) { + for _, test := range []struct { + typ, value string + unit timeseries.FieldDataType + want int64 + valid bool + }{ + {"TimestampSecond", "-1", 0, -1e9, true}, + {"TimestampMillisecond", "-1", 0, -1e6, true}, + {"TimestampMicrosecond", "-1", 0, -1e3, true}, + {"TimestampNanosecond", "1704067200123456789", 0, 1704067200123456789, true}, + {"TimestampSecond", "9223372037", 0, 0, false}, + {"TimestampSecond", "-9223372037", 0, 0, false}, + {"Int64", "1704067200", timeseries.DateTimeUnixSecs, 1704067200000000000, true}, + {"Float64", "1704067200.123456789", timeseries.DateTimeUnixSecs, 1704067200123456789, true}, + {"Float64", "0.0000000001", timeseries.DateTimeUnixSecs, 0, false}, + {"Float64", "1e3", timeseries.DateTimeUnixMilli, 1e9, true}, + {"String", "0", 0, 0, false}, + } { + t.Run(test.typ+test.value, func(t *testing.T) { + typ, _ := fieldType(test.typ) + field := timeseries.FieldDefinition{DataType: typ, SDataType: test.typ, ProviderData1: byte(test.unit)} + got, err := parseEpoch(json.Number(test.value), field) + if (err == nil) != test.valid || (test.valid && got != test.want) { + t.Fatalf("epoch=%d err=%v", got, err) + } + if test.valid { + v, err := epochValue(got, field) + if err != nil { + t.Fatal(err) + } + raw, err := json.Marshal(v) + if err != nil { + t.Fatal(err) + } + back, err := parseEpoch(json.Number(raw), field) + if err != nil || back != got { + t.Fatalf("axis changed: %s %v", raw, err) + } + } + }) + } +} + +type brokenWriter struct{} + +func (brokenWriter) Write([]byte) (int, error) { return 0, io.ErrClosedPipe } + +func TestModelGuardsAndWriter(t *testing.T) { + trq := query() + ts, err := UnmarshalTimeseries([]byte(envelope(`[]`, 0)), trq) + if err != nil { + t.Fatal(err) + } + w := httptest.NewRecorder() + if err := MarshalTimeseriesWriter(ts, nil, 200, w); err != nil { + t.Fatal(err) + } + if w.Header().Get("Content-Type") != "application/json" || w.Header().Get("X-Greptime-Execution-Time") != "0" { + t.Fatal("missing response metadata") + } + if err := MarshalTimeseriesWriter(ts, nil, 200, brokenWriter{}); !errors.Is(err, io.ErrClosedPipe) { + t.Fatal(err) + } + if _, err := UnmarshalTimeseriesReader(nil, trq); err == nil { + t.Fatal("nil reader accepted") + } + if _, err := UnmarshalTimeseries(nil, nil); err == nil { + t.Fatal("nil query accepted") + } + if _, err := UnmarshalTimeseries([]byte(envelope(`[]`, 0)), ×eries.TimeRangeQuery{}); err == nil { + t.Fatal("missing plan accepted") + } + if _, err := MarshalTimeseries(&dataset.DataSet{}, nil, 200); err == nil { + t.Fatal("unknown model accepted") + } + if _, err := marshalCache(nil, nil, 200); err == nil { + t.Fatal("nil cache object accepted") + } + buf, err := marshalCache(ts, nil, 200) + if err != nil { + t.Fatal(err) + } + for _, bad := range [][]byte{nil, []byte("GSQL"), append(bytes.Clone(buf), 0), buf[:len(buf)-1]} { + if _, err := unmarshalCache(bad, trq); err == nil { + t.Fatal("invalid cache accepted") + } + } +} + +func BenchmarkModelRoundTrip(b *testing.B) { + input := []byte(envelope(`[[18446744073709551615,1000000000,"a"],[9007199254740993,2000000000,"b"]]`, 2)) + trq := query() + for b.Loop() { + ts, err := UnmarshalTimeseries(input, trq) + if err != nil { + b.Fatal(err) + } + if _, err := MarshalTimeseries(ts, nil, 200); err != nil { + b.Fatal(err) + } + } +} + +func FuzzModelDecode(f *testing.F) { + f.Add([]byte(envelope(`[[7,1000000000,"a"]]`, 1))) + f.Add([]byte(envelope(`[]`, 0))) + f.Fuzz(func(t *testing.T, body []byte) { + ts, err := UnmarshalTimeseries(body, query()) + if err != nil { + return + } + if _, err := MarshalTimeseries(ts, nil, 200); err != nil { + t.Fatal(err) + } + }) +} diff --git a/pkg/backends/greptimedb/model/stream_test.go b/pkg/backends/greptimedb/model/stream_test.go new file mode 100644 index 000000000..d42607800 --- /dev/null +++ b/pkg/backends/greptimedb/model/stream_test.go @@ -0,0 +1,73 @@ +/* + * Copyright 2026 The Trickster Authors + * 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 model + +import ( + "encoding/json" + "fmt" + "strings" + "testing" + + "github.com/trickstercache/trickster/v2/pkg/timeseries" + "github.com/trickstercache/trickster/v2/pkg/timeseries/dataset" + "github.com/trickstercache/trickster/v2/pkg/timeseries/dataset/stream/streamtest" +) + +func TestSQLStreamConformance(t *testing.T) { + base := envelope(`[[18446744073709551615,2000000000,"a"],[9007199254740993,1000000000,"\u0061"],[null,1000000000,null],[7,2000000000,"null"]]`, 4) + var fields map[string]json.RawMessage + raw, err := json.Marshal(decoded(t, []byte(base)).Output[0].Records) + if err != nil { + t.Fatal(err) + } + if err := json.Unmarshal(raw, &fields); err != nil { + t.Fatal(err) + } + lateSchema := fmt.Sprintf(`{"execution_time_ms":0,"output":[{"records":{"rows":%s,"total_rows":4,"schema":%s,"metrics":{}}}]}`, fields["rows"], fields["schema"]) + for name, body := range map[string]string{ + "rows": base, + "schema_last": lateSchema, + "empty": envelope(`[]`, 0), + "metrics_null": strings.Replace(base, `"total_rows":4`, `"total_rows":4,"metrics":null`, 1), + } { + t.Run(name, func(t *testing.T) { + streamtest.Conformance(t, newDecoder, streamtest.Case{ + TRQ: query(), Body: []byte(body), + Unwrap: func(ts timeseries.Timeseries) *dataset.DataSet { return ts.(*dataSet).DataSet }, + }) + ts, err := UnmarshalTimeseriesReader(strings.NewReader(body), query()) + if err != nil { + t.Fatal(err) + } + if name != "empty" && (ts.SeriesCount() != 3 || ts.ValueCount() != 4) { + t.Fatalf("tag encoding changed identity: %d series, %d values", ts.SeriesCount(), ts.ValueCount()) + } + }) + } + for name, body := range map[string]string{ + "truncated": base[:len(base)-1], + "trailing": base + `{}`, + "duplicate_output": strings.Replace(base, `"output":`, `"output":[],"output":`, 1), + "duplicate_rows": strings.Replace(base, `"rows":`, `"rows":[],"rows":`, 1), + "duplicate_schema": strings.Replace(base, `"total_rows":4`, `"total_rows":4,"schema":{}`, 1), + "wrong_total": strings.Replace(base, `"total_rows":4`, `"total_rows":5`, 1), + "bad_tag_after_good_rows": envelope(`[[1,1000000000,"a"],[2,2000000000,42]]`, 2), + "escaped_duplicate_point": envelope(`[[1,1000000000,"a"],[2,1000000000,"\u0061"]]`, 2), + "schema_last_null_rows": strings.Replace(lateSchema, string(fields["rows"]), "null", 1), + } { + t.Run(name, func(t *testing.T) { + streamtest.Conformance(t, newDecoder, streamtest.Case{TRQ: query(), Body: []byte(body), WantErr: streamtest.ErrAny}) + }) + } +} diff --git a/pkg/backends/greptimedb/model/unmarshal.go b/pkg/backends/greptimedb/model/unmarshal.go new file mode 100644 index 000000000..3bb35bd3e --- /dev/null +++ b/pkg/backends/greptimedb/model/unmarshal.go @@ -0,0 +1,358 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 model + +import ( + "bytes" + "encoding/json" + "errors" + "io" + "slices" + + "github.com/trickstercache/trickster/v2/pkg/parsing/sqlanalyzer" + "github.com/trickstercache/trickster/v2/pkg/timeseries" + "github.com/trickstercache/trickster/v2/pkg/timeseries/dataset" + "github.com/trickstercache/trickster/v2/pkg/timeseries/dataset/stream" + "github.com/trickstercache/trickster/v2/pkg/timeseries/epoch" +) + +type column struct { + Name string `json:"name"` + Type string `json:"data_type"` +} + +type schema struct { + Columns []column `json:"column_schemas"` +} + +type records struct { + Schema schema `json:"schema"` + Rows [][]any `json:"rows"` + Total *uint64 `json:"total_rows"` + Metrics map[string]json.RawMessage `json:"metrics,omitempty"` +} + +type output struct { + Records *records `json:"records"` +} + +type response struct { + Output []output `json:"output"` + ExecutionTime *uint64 `json:"execution_time_ms"` +} + +func UnmarshalTimeseries(body []byte, trq *timeseries.TimeRangeQuery) (timeseries.Timeseries, error) { + return stream.BytesUnmarshaler(newDecoder)(body, trq) +} + +func UnmarshalTimeseriesReader(reader io.Reader, trq *timeseries.TimeRangeQuery) (timeseries.Timeseries, error) { + return stream.ReaderUnmarshaler(newDecoder)(reader, trq) +} + +type sqlDecoder struct { + trq *timeseries.TimeRangeQuery + plan *sqlanalyzer.QueryPlan + fields timeseries.FieldDefinitions + orderedFloats []bool + builder *dataset.Builder + row []json.RawMessage + rowCount uint64 + tagErr error +} + +func newDecoder(trq *timeseries.TimeRangeQuery) (stream.Decoder, error) { + if trq == nil { + return nil, timeseries.ErrInvalidBody + } + plan, ok := trq.ParsedQuery.(*sqlanalyzer.QueryPlan) + if !ok || plan == nil || trq.Step <= 0 { + return nil, timeseries.ErrInvalidBody + } + d := &sqlDecoder{trq: trq, plan: plan} + return stream.NewJSON(d.walk, func() (timeseries.Timeseries, error) { + ds, err := d.builder.Finish() + if err != nil { + return nil, err + } + return &dataSet{DataSet: ds, fields: d.fields}, nil + }), nil +} + +func (d *sqlDecoder) walk(dec *json.Decoder) error { + dec.DisallowUnknownFields() + var hasOutput bool + var elapsed *uint64 + err := stream.Object(dec, func(key string) error { + switch key { + case "output": + if hasOutput { + return timeseries.ErrInvalidBody + } + hasOutput = true + count := 0 + err := stream.Array(dec, func() error { + count++ + if count != 1 { + return timeseries.ErrInvalidBody + } + return stream.Object(dec, func(key string) error { + if key != "records" || d.builder != nil { + return timeseries.ErrInvalidBody + } + return d.readRecords(dec) + }) + }) + if err == nil && count != 1 { + return timeseries.ErrInvalidBody + } + return err + case "execution_time_ms": + if elapsed != nil { + return timeseries.ErrInvalidBody + } + return dec.Decode(&elapsed) + } + return timeseries.ErrInvalidBody + }) + if err == nil && (!hasOutput || elapsed == nil || d.builder == nil) { + return timeseries.ErrInvalidBody + } + return err +} + +func (d *sqlDecoder) readRecords(dec *json.Decoder) error { + var rowsSeen bool + var total *uint64 + var pending json.RawMessage + err := stream.Object(dec, func(key string) error { + switch key { + case "schema": + if d.builder != nil { + return timeseries.ErrInvalidBody + } + var schema schema + if err := dec.Decode(&schema); err != nil { + return err + } + return d.setSchema(schema.Columns) + case "rows": + if rowsSeen { + return timeseries.ErrInvalidBody + } + rowsSeen = true + if d.builder == nil { + // JSON object keys are unordered; only this alternate layout needs a buffer. + return dec.Decode(&pending) + } + return d.readRows(dec) + case "total_rows": + if total != nil { + return timeseries.ErrInvalidBody + } + return dec.Decode(&total) + case "metrics": + err := stream.Object(dec, func(string) error { return timeseries.ErrInvalidBody }) + if errors.Is(err, stream.ErrNull) { + return nil + } + return err + } + return timeseries.ErrInvalidBody + }) + if err != nil { + return err + } + if d.builder == nil || !rowsSeen || total == nil { + return timeseries.ErrInvalidBody + } + if pending != nil { + if err := d.readRows(json.NewDecoder(bytes.NewReader(pending))); err != nil { + return err + } + } + if *total != d.rowCount { + return timeseries.ErrInvalidBody + } + return nil +} + +func (d *sqlDecoder) setSchema(columns []column) error { + fields, err := fieldDefinitions(columns, d.plan) + if err != nil { + return err + } + d.fields = fields + d.orderedFloats = make([]bool, len(fields)) + var sf timeseries.SeriesFields + for i, field := range fields { + d.orderedFloats[i] = field.DataType == timeseries.Float64 && slices.ContainsFunc(d.plan.Ordering, func(term timeseries.OrderTerm) bool { + return term.Column == field.Name + }) + switch field.Role { + case timeseries.RoleTimestamp: + sf.Timestamp = field + case timeseries.RoleTag: + sf.Tags = append(sf.Tags, field) + case timeseries.RoleValue: + sf.Values = append(sf.Values, field) + } + } + d.builder = dataset.NewBuilder(d.trq, dataset.BuilderOptions{ + Fields: sf, SeriesName: "sql", QueryStatement: d.trq.Statement, + Duplicates: dataset.DuplicatesError, TagString: d.tagString, + }) + return nil +} + +func (d *sqlDecoder) tagString(field timeseries.FieldDefinition, raw []byte) string { + value, err := decodeRawValue(raw, field.SDataType) + if err != nil { + d.tagErr = err + return "" + } + encoded, err := json.Marshal(value) + if err != nil { + d.tagErr = err + } + return string(encoded) +} + +func (d *sqlDecoder) readRows(dec *json.Decoder) error { + return stream.Array(dec, func() error { + if err := dec.Decode(&d.row); err != nil { + return err + } + if len(d.row) != len(d.fields) { + return timeseries.ErrInvalidBody + } + row := d.builder.Row() + tag := 0 + for i, field := range d.fields { + raw := d.row[i] + // JSON null cannot distinguish SQL NULL from NaN/Inf for numeric sorting. + if d.orderedFloats[i] && bytes.Equal(raw, []byte("null")) { + return timeseries.ErrInvalidBody + } + if field.Role == timeseries.RoleTag { + row.SetTag(tag, raw) + tag++ + continue + } + value, err := decodeRawValue(raw, field.SDataType) + if err != nil { + return err + } + if field.Role == timeseries.RoleTimestamp { + ep, err := parseEpoch(json.Number(raw), field) + if err != nil || !onGrid(ep, d.trq) { + return timeseries.ErrInvalidTimeFormat + } + row.SetEpoch(epoch.Epoch(ep)) + } else { + row.AddValue(value) + } + } + if err := row.Commit(); err != nil { + return err + } + d.rowCount++ + return d.tagErr + }) +} + +func decodeRawValue(raw []byte, typ string) (any, error) { + if bytes.Equal(raw, []byte("null")) { + return nil, nil + } + switch typ { + case "String": + var value string + err := json.Unmarshal(raw, &value) + return value, err + case "Boolean": + var value bool + err := json.Unmarshal(raw, &value) + return value, err + } + return decodeValue(json.Number(raw), typ) +} + +func fieldDefinitions(columns []column, plan *sqlanalyzer.QueryPlan) (timeseries.FieldDefinitions, error) { + roles := map[string]timeseries.FieldRole{plan.OutputColumn: timeseries.RoleTimestamp} + for _, name := range plan.GroupColumns { + roles[name] = timeseries.RoleTag + } + for _, name := range plan.ValueColumns { + roles[name] = timeseries.RoleValue + } + if len(roles) > len(columns) || (len(plan.ValueColumns) > 0 && len(roles) != len(columns)) { + return nil, timeseries.ErrInvalidBody + } + fields := make(timeseries.FieldDefinitions, len(columns)) + seen := make(map[string]bool, len(columns)) + values := 0 + for i, col := range columns { + role, ok := roles[col.Name] + // The Cockroach adapter currently names only timestamp/group fields. + // Remaining columns are typed values supplied by the origin schema. + if !ok && len(plan.ValueColumns) == 0 { + role, ok = timeseries.RoleValue, true + } + typ, supported := fieldType(col.Type) + if !ok || !supported || col.Name == "" || seen[col.Name] { + return nil, timeseries.ErrInvalidBody + } + seen[col.Name] = true + delete(roles, col.Name) + field := timeseries.FieldDefinition{ + Name: col.Name, SDataType: col.Type, + DataType: typ, Role: role, OutputPosition: i, + } + if role == timeseries.RoleTimestamp { + field.ProviderData1 = byte(plan.OutputUnit) + if axisScale(field) == 0 { + return nil, timeseries.ErrInvalidTimeFormat + } + } + fields[i] = field + if role == timeseries.RoleValue { + values++ + } + } + // The common dataset treats a series without values as empty when cropping. + if len(roles) != 0 || values == 0 { + return nil, timeseries.ErrInvalidBody + } + for _, term := range plan.Ordering { + if !seen[term.Column] { + return nil, timeseries.ErrInvalidBody + } + } + return fields, nil +} + +func onGrid(value int64, trq *timeseries.TimeRangeQuery) bool { + step := int64(trq.Step) + r, phase := value%step, int64(trq.Phase)%step + if r < 0 { + r += step + } + if phase < 0 { + phase += step + } + return r == phase +} diff --git a/pkg/backends/greptimedb/model/value.go b/pkg/backends/greptimedb/model/value.go new file mode 100644 index 000000000..fe28446ec --- /dev/null +++ b/pkg/backends/greptimedb/model/value.go @@ -0,0 +1,167 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 model + +import ( + "encoding/json" + "math" + "math/big" + "strconv" + "strings" + + "github.com/trickstercache/trickster/v2/pkg/timeseries" +) + +func fieldType(name string) (timeseries.FieldDataType, bool) { + switch name { + case "Int8", "Int16", "Int32", "Int64", "TimestampSecond", "TimestampMillisecond", "TimestampMicrosecond", "TimestampNanosecond", "Date": + return timeseries.Int64, true + case "UInt8", "UInt16", "UInt32", "UInt64": + return timeseries.Uint64, true + case "Float32", "Float64": + return timeseries.Float64, true + case "String": + return timeseries.String, true + case "Boolean": + return timeseries.Bool, true + case "Null": + return timeseries.Null, true + } + return timeseries.Unknown, false +} + +func decodeValue(value any, typ string) (any, error) { + if value == nil { + return nil, nil + } + kind, ok := fieldType(typ) + if !ok { + return nil, timeseries.ErrInvalidBody + } + switch kind { + case timeseries.String: + if v, ok := value.(string); ok { + return v, nil + } + case timeseries.Bool: + if v, ok := value.(bool); ok { + return v, nil + } + case timeseries.Int64, timeseries.Uint64, timeseries.Float64: + n, ok := value.(json.Number) + if !ok { + return nil, timeseries.ErrInvalidBody + } + bits := 64 + for _, width := range []int{8, 16, 32} { + if typ == "Int"+strconv.Itoa(width) || typ == "UInt"+strconv.Itoa(width) || typ == "Float"+strconv.Itoa(width) { + bits = width + } + } + switch kind { + case timeseries.Int64: + return strconv.ParseInt(string(n), 10, bits) + case timeseries.Uint64: + return strconv.ParseUint(string(n), 10, bits) + case timeseries.Float64: + // JSON carries a decimal rendering of a Float32; do not widen a + // rounded binary32 and change that decimal during reconstruction. + v, err := strconv.ParseFloat(string(n), 64) + if err == nil && !math.IsNaN(v) && !math.IsInf(v, 0) && (bits == 64 || math.Abs(v) <= math.MaxFloat32) { + return v, nil + } + } + } + return nil, timeseries.ErrInvalidBody +} + +func axisScale(field timeseries.FieldDefinition) int64 { + switch field.SDataType { + case "TimestampSecond": + return 1e9 + case "TimestampMillisecond": + return 1e6 + case "TimestampMicrosecond": + return 1e3 + case "TimestampNanosecond": + return 1 + } + if !strings.HasPrefix(field.SDataType, "Int") && !strings.HasPrefix(field.SDataType, "UInt") && + field.SDataType != "Float64" && field.SDataType != "Float32" { + return 0 + } + switch timeseries.FieldDataType(field.ProviderData1) { + case timeseries.DateTimeUnixSecs: + return 1e9 + case timeseries.DateTimeUnixMilli: + return 1e6 + case timeseries.DateTimeUnixMicro: + return 1e3 + case timeseries.DateTimeUnixNano: + return 1 + } + return 0 +} + +func parseEpoch(value any, field timeseries.FieldDefinition) (int64, error) { + n, ok := value.(json.Number) + if !ok { + return 0, timeseries.ErrInvalidTimeFormat + } + scale := axisScale(field) + if scale == 0 { + return 0, timeseries.ErrInvalidTimeFormat + } + if field.DataType == timeseries.Float64 { + r, ok := new(big.Rat).SetString(string(n)) + if !ok { + return 0, timeseries.ErrInvalidTimeFormat + } + r.Mul(r, new(big.Rat).SetInt64(scale)) + if !r.IsInt() || !r.Num().IsInt64() { + return 0, timeseries.ErrInvalidTimeFormat + } + return r.Num().Int64(), nil + } + v, err := strconv.ParseInt(string(n), 10, 64) + if err != nil || v > math.MaxInt64/scale || v < math.MinInt64/scale { + return 0, timeseries.ErrInvalidTimeFormat + } + return v * scale, nil +} + +func epochValue(ep int64, field timeseries.FieldDefinition) (any, error) { + scale := axisScale(field) + if scale == 0 { + return nil, timeseries.ErrInvalidTimeFormat + } + if field.DataType == timeseries.Float64 { + r := new(big.Rat).SetFrac(big.NewInt(ep), big.NewInt(scale)) + return json.Number(r.FloatString(9)), nil + } + if ep%scale != 0 { + return nil, timeseries.ErrInvalidTimeFormat + } + value := ep / scale + if field.DataType == timeseries.Uint64 { + if value < 0 { + return nil, timeseries.ErrInvalidTimeFormat + } + return uint64(value), nil + } + return value, nil +} diff --git a/pkg/backends/greptimedb/mysql.go b/pkg/backends/greptimedb/mysql.go new file mode 100644 index 000000000..a27937c0c --- /dev/null +++ b/pkg/backends/greptimedb/mysql.go @@ -0,0 +1,185 @@ +/* + * Copyright 2026 The Trickster Authors + * 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 greptimedb + +import ( + "errors" + "strings" + "time" + + "github.com/trickstercache/trickster/v2/pkg/backends/mysql" + "github.com/trickstercache/trickster/v2/pkg/backends/providers" + "github.com/trickstercache/trickster/v2/pkg/parsing/sqlanalyzer" + "github.com/trickstercache/trickster/v2/pkg/parsing/sqlanalyzer/cockroach" + "github.com/trickstercache/trickster/v2/pkg/parsing/sqlanalyzer/vitess" + "github.com/trickstercache/trickster/v2/pkg/timeseries" + + vtmysql "vitess.io/vitess/go/mysql" + "vitess.io/vitess/go/sqltypes" + "vitess.io/vitess/go/vt/sqlparser" +) + +type mysqlEngine struct{} + +var mysqlAnalyzer = newMySQLAnalyzer() + +// MySQLEngine returns Greptime's MySQL wire-protocol implementation. +func MySQLEngine() mysql.Engine { return mysqlEngine{} } + +func (mysqlEngine) Name() string { return providers.GreptimeDB } +func (mysqlEngine) DefaultPort() string { return "4002" } +func (mysqlEngine) SupportsHTTP() bool { return true } +func (mysqlEngine) Analyzer(view mysql.SessionView) sqlanalyzer.DialectAnalyzer { + return &mysqlDialectAnalyzer{inner: mysqlAnalyzer, utc: view.TimeZone == "UTC" || view.TimeZone == "+00:00"} +} + +func (mysqlEngine) StreamState(conn *vtmysql.Conn) (uint16, uint16, error) { + if conn == nil { + return 0, 0, errors.New("nil GreptimeDB MySQL connection") + } + // Greptime's record-batch writer ends resultsets with default EOF status + // and warning fields. Its diagnostic SQL is different from MySQL's and + // must not run here: even SHOW COUNT(*) WARNINGS is unsupported. Statements + // returning OK packets retain their actual metadata in the shared path. + if result := conn.StreamOKResult(); result != nil { + return result.StatusFlags, 0, nil + } + return 0, 0, nil +} + +func (mysqlEngine) InitSession(conn *vtmysql.Conn) (mysql.SessionView, error) { + if conn == nil { + return mysql.SessionView{}, errors.New("nil GreptimeDB MySQL connection") + } + r, err := conn.ExecuteFetch("SHOW TIMEZONE", 1, true) + if err != nil { + return mysql.SessionView{}, err + } + if r == nil || len(r.Rows) != 1 || len(r.Rows[0]) != 1 || r.Rows[0][0].IsNull() { + return mysql.SessionView{}, errors.New("invalid GreptimeDB time zone") + } + return mysql.SessionView{TimeZone: r.Rows[0][0].ToString()}, nil +} + +func (mysqlEngine) ResultSemantics() mysql.ResultSemantics { + return mysql.ResultSemantics{Timestamp: mysqlTimestampEpoch, BinaryText: true, NullsLast: true} +} + +func mysqlTimestampEpoch(value sqltypes.Value) (int64, error) { + t, err := time.Parse("2006-01-02 15:04:05.999999999", value.ToString()) + if err != nil { + return 0, err + } + if !sqlanalyzer.SafeUnixSeconds(t.Unix()) { + return 0, errors.New("GreptimeDB timestamp is outside the cache range") + } + return t.UnixNano(), nil +} + +func (c *Client) MySQLRouteConfig() (mysql.ProtocolConfig, error) { + return mysql.ProtocolConfigForEngine(c.Configuration(), MySQLEngine()) +} + +type mysqlDialectAnalyzer struct { + inner *vitess.Analyzer + utc bool +} + +func newMySQLAnalyzer() *vitess.Analyzer { + a, err := vitess.NewAnalyzerWithOptions(vitess.Options{ + BucketMatchers: []vitess.BucketMatcher{mysqlTimestampBucket}, + DeterministicFunctions: []string{"date_bin", "date_trunc"}, + }) + if err != nil { + panic(err) + } + return a +} + +func (a *mysqlDialectAnalyzer) Analyze(query string, _ time.Time) sqlanalyzer.Analysis { + stmt, err := a.inner.Parser().Parse(query) + return a.AnalyzeParsed(query, stmt, err) +} + +func (a *mysqlDialectAnalyzer) AnalyzeParsed(query string, stmt sqlparser.Statement, err error) sqlanalyzer.Analysis { + analysis := a.inner.AnalyzeParsed(query, stmt, err) + if analysis.Mode != sqlanalyzer.CacheModeDelta { + return analysis + } + p := analysis.Plan + // MySQL's integer division, epoch inference and implicit casts are not + // Greptime contracts. Timestamp bucket results are lossless in UTC only. + if !a.utc || p.OutputUnit != timeseries.DateTimeSQL || p.InputUnit != timeseries.DateTimeSQL || + p.Step%time.Second != 0 { + return sqlanalyzer.ObjectAnalysis(sqlanalyzer.ReasonUnsupportedBucket, errUnrenderable) + } + unsafe := false + _ = sqlparser.Walk(func(node sqlparser.SQLNode) (bool, error) { + switch n := node.(type) { + case *sqlparser.BetweenExpr: + unsafe = true + case *sqlparser.ComparisonExpr: + if n.Operator == sqlparser.LessEqualOp { + unsafe = true + } + } + return !unsafe, nil + }, stmt) + if unsafe { + return sqlanalyzer.ObjectAnalysis(sqlanalyzer.ReasonUnsafePredicate, errUnrenderable) + } + // Cache only complete buckets, matching the provider's HTTP and PGWire paths. + lower := sqlanalyzer.CeilBucket(p.LowerBound.Value, p.Step, p.Phase) + upper := sqlanalyzer.FloorBucket(p.UpperBound.Value, p.Step, p.Phase) + if !upper.After(lower) { + return sqlanalyzer.ObjectAnalysis(sqlanalyzer.ReasonUnsafePredicate, errUnrenderable) + } + p.DropsPartialBuckets = !lower.Equal(p.LowerBound.Value) || !upper.Equal(p.UpperBound.Value) + p.LowerBound.Value, p.UpperBound.Value = lower, upper + return analysis +} + +func mysqlTimestampBucket(expr sqlparser.Expr) (*sqlparser.ColName, time.Duration, timeseries.FieldDataType, bool) { + fn, ok := expr.(*sqlparser.FuncExpr) + if !ok || !fn.Qualifier.IsEmpty() { + return nil, 0, 0, false + } + name := strings.ToLower(fn.Name.String()) + if (name != "date_bin" || len(fn.Exprs) != 3) && (name != "date_trunc" || len(fn.Exprs) != 2) { + return nil, 0, 0, false + } + width, ok := fn.Exprs[0].(*sqlparser.Literal) + if !ok || width.Type != sqlparser.StrVal { + return nil, 0, 0, false + } + column, ok := fn.Exprs[1].(*sqlparser.ColName) + if !ok { + return nil, 0, 0, false + } + var step time.Duration + if name == "date_bin" { + anchor, ok := fn.Exprs[2].(*sqlparser.FuncExpr) + if !ok || !anchor.Qualifier.IsEmpty() || !strings.EqualFold(anchor.Name.String(), "from_unixtime") || len(anchor.Exprs) != 1 { + return nil, 0, 0, false + } + zero, ok := anchor.Exprs[0].(*sqlparser.Literal) + if !ok || zero.Type != sqlparser.IntVal || zero.Val != "0" { + return nil, 0, 0, false + } + step, _ = cockroach.ParseCompactDuration(width.Val, compactUnits) + } else { + step = map[string]time.Duration{"second": time.Second, "minute": time.Minute, "hour": time.Hour, "day": 24 * time.Hour}[strings.ToLower(width.Val)] + } + return column, step, timeseries.DateTimeSQL, step > 0 +} diff --git a/pkg/backends/greptimedb/mysql_test.go b/pkg/backends/greptimedb/mysql_test.go new file mode 100644 index 000000000..e16afcc11 --- /dev/null +++ b/pkg/backends/greptimedb/mysql_test.go @@ -0,0 +1,93 @@ +/* + * Copyright 2026 The Trickster Authors + * 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 greptimedb + +import ( + "strings" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/backends/mysql" + "github.com/trickstercache/trickster/v2/pkg/parsing/sqlanalyzer" + "github.com/trickstercache/trickster/v2/pkg/timeseries" + + "vitess.io/vitess/go/sqltypes" + querypb "vitess.io/vitess/go/vt/proto/query" +) + +const mysqlBucketQuery = "SELECT DATE_BIN('1m', ts, FROM_UNIXTIME(0)) AS time, label, COUNT(*) AS value " + + "FROM readings WHERE ts >= FROM_UNIXTIME(1767225600) AND ts < FROM_UNIXTIME(1767225720) " + + "GROUP BY time, label ORDER BY time, label" + +func TestMySQLBucketAnalysis(t *testing.T) { + a := MySQLEngine().Analyzer(mysql.SessionView{TimeZone: "UTC"}) + for _, bucket := range []string{"DATE_BIN('1m', ts, FROM_UNIXTIME(0))", "DATE_TRUNC('minute', ts)"} { + t.Run(bucket, func(t *testing.T) { + query := strings.Replace(mysqlBucketQuery, "DATE_BIN('1m', ts, FROM_UNIXTIME(0))", bucket, 1) + analysis := a.Analyze(query, time.Time{}) + if analysis.Mode != sqlanalyzer.CacheModeDelta || analysis.Plan == nil { + t.Fatalf("not delta cacheable: %+v", analysis) + } + if analysis.Plan.OutputUnit != timeseries.DateTimeSQL || analysis.Plan.Step != time.Minute { + t.Fatal("wrong time axis") + } + rendered, err := analysis.Plan.RenderExtent(timeseries.Extent{Start: time.Unix(1767225660, 0), End: time.Unix(1767225720, 0)}) + if err != nil || !strings.Contains(strings.ToLower(rendered), "from_unixtime(1767225780)") { + t.Fatalf("extent = %s, %v", rendered, err) + } + }) + } + for _, query := range []string{ + strings.Replace(mysqlBucketQuery, "'1m'", "'1M'", 1), + strings.Replace(mysqlBucketQuery, "'1m'", "'0m'", 1), + strings.Replace(mysqlBucketQuery, "'1m'", "'+1m'", 1), + strings.Replace(mysqlBucketQuery, "'1m'", "'9223372036854775807m'", 1), + strings.Replace(mysqlBucketQuery, "'1m'", "'500ms'", 1), + strings.Replace(mysqlBucketQuery, "'1m'", "'1500ms'", 1), + strings.Replace(mysqlBucketQuery, "FROM_UNIXTIME(0)", "FROM_UNIXTIME(1)", 1), + strings.Replace(mysqlBucketQuery, "1767225600", "1767225719", 1), + strings.Replace(mysqlBucketQuery, "ts <", "ts <=", 1), + } { + if got := a.Analyze(query, time.Time{}); got.Mode == sqlanalyzer.CacheModeDelta { + t.Fatalf("unsafe query admitted: %s", query) + } + } + for _, zone := range []string{"", "+08:00", "America/New_York"} { + if got := MySQLEngine().Analyzer(mysql.SessionView{TimeZone: zone}).Analyze(mysqlBucketQuery, time.Time{}); got.Mode == sqlanalyzer.CacheModeDelta { + t.Fatalf("unverified zone admitted: %s", zone) + } + } +} + +func TestMySQLTimestampPrecision(t *testing.T) { + for _, raw := range []string{"2026-01-01 00:00:00", "2026-01-01 00:00:00.123456789", "1969-12-31 23:59:59.999999999"} { + value := sqltypes.MakeTrusted(querypb.Type_TIMESTAMP, []byte(raw)) + got, err := mysqlTimestampEpoch(value) + want, parseErr := time.Parse("2006-01-02 15:04:05.999999999", raw) + if err != nil || parseErr != nil || got != want.UnixNano() { + t.Fatalf("%s: %d %v", raw, got, err) + } + } + for _, raw := range []string{"0000-00-00 00:00:00", "9999-12-31 23:59:59", "not a timestamp"} { + if _, err := mysqlTimestampEpoch(sqltypes.MakeTrusted(querypb.Type_TIMESTAMP, []byte(raw))); err == nil { + t.Fatalf("accepted %q", raw) + } + } + if _, err := (mysqlEngine{}).InitSession(nil); err == nil { + t.Fatal("accepted nil session") + } + if _, _, err := (mysqlEngine{}).StreamState(nil); err == nil { + t.Fatal("accepted nil stream") + } +} diff --git a/pkg/backends/greptimedb/prometheus.go b/pkg/backends/greptimedb/prometheus.go new file mode 100644 index 000000000..2f827994b --- /dev/null +++ b/pkg/backends/greptimedb/prometheus.go @@ -0,0 +1,127 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 greptimedb + +import ( + "mime" + "net/http" + "net/url" + "strings" + + "github.com/trickstercache/trickster/v2/pkg/backends/prometheus" + "github.com/trickstercache/trickster/v2/pkg/proxy/params" + "github.com/trickstercache/trickster/v2/pkg/proxy/request" + + "github.com/prometheus/prometheus/promql/parser" +) + +const ( + promPath = "/v1/prometheus" + databaseParam = "db" + databaseHeader = "X-Greptime-Db-Name" +) + +func promHooks() prometheus.Hooks { + return prometheus.Hooks{ + PathPrefix: promPath, + CacheKeyParams: []string{databaseParam, "lookback"}, + CacheKeyHeaders: []string{databaseHeader, "X-Greptime-Timezone", "X-Greptime-Auth"}, + PrepareRequest: preparePromRequest, + PreserveQueryGrid: true, + AlignQueryGrid: true, + } +} + +func isPromRange(r *http.Request) bool { + return r != nil && r.URL != nil && strings.HasSuffix(r.URL.Path, promPath+prometheus.APIPath+"query_range") +} + +func preparePromRequest(r *http.Request) bool { + if r.Method != http.MethodGet && r.Method != http.MethodPost { + return false + } + if r.Method == http.MethodPost && strings.Contains(r.URL.Path, prometheus.APIPath+"label/") { + return false + } + if rsc := request.GetResources(r); rsc != nil && rsc.PathConfig != nil && len(rsc.PathConfig.RequestParams) != 0 { + return false + } + values, err := url.ParseQuery(r.URL.RawQuery) + if err != nil || !validPromParams(values) { + return false + } + if r.Method == http.MethodPost { + contentType, _, err := mime.ParseMediaType(r.Header.Get("Content-Type")) + if err != nil || contentType != "application/x-www-form-urlencoded" { + return false + } + body, err := request.GetBody(r) + if err != nil { + return false + } + form, err := url.ParseQuery(string(body)) + if err != nil || !validPromParams(form) { + return false + } + // Greptime reads db only from the URL/context, and other fields prefer + // the URL even when its explicit value is empty. + for name, field := range form { + if _, exists := values[name]; !exists && name != databaseParam { + values[name] = field + } + } + } + if query := values.Get("query"); query != "" && !cacheablePromExpression(query) { + return false + } + if r.Method == http.MethodPost { + params.SetRequestValues(r, values) + } + return true +} + +func cacheablePromExpression(query string) bool { + expr, err := metricParser.ParseExpr(query) + if err != nil { + return false + } + valid := true + parser.Inspect(expr, func(node parser.Node, _ []parser.Node) error { + if agg, ok := node.(*parser.AggregateExpr); ok && agg.Op == parser.COUNT_VALUES { + // Greptime currently emits numeric count_values labels as value + // fields, losing series identity and producing duplicate timestamps. + valid = false + } + return nil + }) + return valid +} + +func validPromParams(values url.Values) bool { + for name, fields := range values { + switch name { + case "match[]": + continue + case "query", "start", "end", "step", "time", databaseParam, "lookback", "stats": + if len(fields) == 1 { + continue + } + } + return false + } + return true +} diff --git a/pkg/backends/greptimedb/prometheus_test.go b/pkg/backends/greptimedb/prometheus_test.go new file mode 100644 index 000000000..f049111cf --- /dev/null +++ b/pkg/backends/greptimedb/prometheus_test.go @@ -0,0 +1,449 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 greptimedb + +import ( + "bytes" + "encoding/json" + "fmt" + "io" + "net/http" + "net/http/httptest" + "net/url" + "slices" + "strconv" + "strings" + "sync/atomic" + "testing" + "time" + + bo "github.com/trickstercache/trickster/v2/pkg/backends/options" + "github.com/trickstercache/trickster/v2/pkg/backends/providers" + ep "github.com/trickstercache/trickster/v2/pkg/encoding/providers" + "github.com/trickstercache/trickster/v2/pkg/proxy/headers" + "github.com/trickstercache/trickster/v2/pkg/proxy/request" + "github.com/trickstercache/trickster/v2/pkg/timeseries" + "github.com/trickstercache/trickster/v2/pkg/timeseries/dataset" + "github.com/trickstercache/trickster/v2/pkg/timeseries/merge" +) + +func TestPrometheusPaths(t *testing.T) { + o := bo.New() + b, err := NewClient("test", o, nil, nil, nil, nil) + if err != nil { + t.Fatal(err) + } + c := b.(*Client) + paths := c.DefaultPathConfigs(o) + for _, name := range []string{"query", "query_range", "series", "labels", "label/job/values"} { + p := paths.Match("GET", promPath+"/api/v1/"+name) + if p == nil || p.HandlerName == "proxy" || !slices.Contains(p.CacheKeyParams, databaseParam) || !slices.Contains(p.CacheKeyHeaders, databaseHeader) { + t.Fatalf("missing route identity: %s %+v", name, p) + } + } + for _, path := range []string{"/api/v1/query", promPath + "/api/v1/alerts", promPath + "/write", "/v1/loki/api/v1/push"} { + if p := paths.Match("POST", path); p == nil || p.HandlerName != "proxy" { + t.Fatalf("unsupported endpoint must proxy: %s", path) + } + } + if len(c.MergeablePaths()) != 5 || !providers.IsPrometheusCompatible("greptimedb") || o.FastForwardPath.Path != promPath+"/api/v1/query" { + t.Fatal("missing prefixed fast-forward/merge capability") + } + o.FastForwardPath.CacheKeyParams[0] = "changed" + if paths.Match("GET", promPath+"/api/v1/query").CacheKeyParams[0] == "changed" { + t.Fatal("fast-forward path shares mutable identity") + } +} + +func TestPrometheusParameterPrecedence(t *testing.T) { + for _, tc := range []struct { + name, query, form, wantQuery, wantDB string + }{ + {"url_wins", "query=vector(1)&db=url", "query=vector(2)&db=form&step=15", "vector(1)", "url"}, + {"empty_url_wins", "query=&db=", "query=vector(2)&db=form", "", ""}, + {"form_db_is_ignored", "", "query=vector(2)&db=form", "vector(2)", ""}, + } { + t.Run(tc.name, func(t *testing.T) { + r := httptest.NewRequest("POST", promPath+"/api/v1/query?"+tc.query, strings.NewReader(tc.form)) + r.Header.Set("Content-Type", "application/x-www-form-urlencoded") + if !preparePromRequest(r) { + t.Fatal("valid request rejected") + } + body, err := io.ReadAll(r.Body) + if err != nil { + t.Fatal(err) + } + v, err := url.ParseQuery(string(body)) + if err != nil || v.Get("query") != tc.wantQuery || v.Get("db") != tc.wantDB || r.URL.RawQuery != string(body) { + t.Fatalf("origin semantics changed: %s %s", r.URL.RawQuery, body) + } + }) + } + for _, raw := range []string{"query=a&query=b", "query=%ZZ", "query=a&limit=1"} { + r := httptest.NewRequest("POST", promPath+"/api/v1/query", strings.NewReader(raw)) + r.Header.Set("Content-Type", "application/x-www-form-urlencoded") + if preparePromRequest(r) { + t.Fatalf("unsupported request was normalized: %s", raw) + } + body, _ := io.ReadAll(r.Body) + if string(body) != raw { + t.Fatal("passthrough lost the original body") + } + } +} + +func (h *httpHarness) promQuery(t *testing.T, method, endpoint string, v url.Values, hdr http.Header) *httptest.ResponseRecorder { + t.Helper() + path := promPath + "/api/v1/" + endpoint + var body io.Reader + if method == "POST" { + body = strings.NewReader(v.Encode()) + // Greptime selects databases only from the URL or header. + if _, exists := v["db"]; exists { + path += "?db=" + url.QueryEscape(v.Get("db")) + } + } else { + path += "?" + v.Encode() + } + r := httptest.NewRequest(method, path, body) + if hdr != nil { + r.Header = hdr.Clone() + } + if method == "POST" { + r.Header.Set("Content-Type", "application/x-www-form-urlencoded") + } + res := h.resources + pc := h.client.DefaultPathConfigs(res.BackendOptions).Match(method, r.URL.Path) + r = request.SetResources(r, request.NewResources(res.BackendOptions, pc, res.CacheConfig, res.CacheClient, h.client, res.Tracer)) + w := httptest.NewRecorder() + h.client.Handlers()[pc.HandlerName].ServeHTTP(w, r) + return w +} + +func TestPrometheusCacheFlow(t *testing.T) { + for _, method := range []string{"GET", "POST"} { + for _, step := range []time.Duration{15 * time.Second, 500 * time.Millisecond} { + t.Run(method+step.String(), func(t *testing.T) { + var calls atomic.Int32 + origin := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + calls.Add(1) + if err := r.ParseForm(); err != nil { + t.Error(err) + } + start, err := time.Parse(time.RFC3339Nano, r.Form.Get("start")) + if err != nil { + t.Error(err) + } + end, err := time.Parse(time.RFC3339Nano, r.Form.Get("end")) + if err != nil { + t.Error(err) + } + values := make([][]any, 0) + for at := start; !at.After(end); at = at.Add(step) { + values = append(values, []any{float64(at.UnixMilli()) / 1000, "2"}) + } + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(map[string]any{"status": "success", "data": map[string]any{"resultType": "matrix", "result": []any{map[string]any{"metric": map[string]string{"job": "test"}, "values": values}}}}) + }) + h := newHTTPHarness(t, origin) + h.resources.BackendOptions.FastForwardDisable = true + start := time.Now().UTC().Add(-5 * time.Minute).Truncate(step).Add(125 * time.Millisecond) + v := url.Values{"query": {"up"}, "db": {"public"}, "start": {start.Format(time.RFC3339Nano)}, "step": {strconv.FormatFloat(step.Seconds(), 'f', -1, 64)}} + for i, want := range []string{"kmiss", "hit", "phit", "hit"} { + points := 3 + if i > 1 { + points = 5 + } + v.Set("end", start.Add(time.Duration(points-1)*step).Format(time.RFC3339Nano)) + w := h.promQuery(t, method, "query_range", v, nil) + engine, status := headers.ParseResultEngineStatus(w.Header().Get(headers.NameTricksterResult)) + if w.Code != 200 || engine != "DeltaProxyCache" || status != want { + t.Fatalf("%d %s %s %s", w.Code, engine, status, w.Body.String()) + } + var got struct { + Data struct{ Result []struct{ Values [][]any } } + } + if err := json.Unmarshal(w.Body.Bytes(), &got); err != nil || len(got.Data.Result) != 1 || len(got.Data.Result[0].Values) != points { + t.Fatalf("wrong matrix: %s (%v)", w.Body.String(), err) + } + for n, point := range got.Data.Result[0].Values { + if point[0] != float64(start.Truncate(step).Add(time.Duration(n)*step).UnixMilli())/1000 || point[1] != "2" { + t.Fatalf("changed evaluation grid: %v", point) + } + } + } + v.Set("start", start.Add(100*time.Millisecond).Format(time.RFC3339Nano)) + w := h.promQuery(t, method, "query_range", v, nil) + engine, status := headers.ParseResultEngineStatus(w.Header().Get(headers.NameTricksterResult)) + if w.Code != 200 || engine != "DeltaProxyCache" || status != "hit" { + t.Fatalf("equivalent aligned range missed cache: %d %s %s", w.Code, engine, status) + } + if calls.Load() != 2 { + t.Fatalf("expected initial and missing extent only: %d", calls.Load()) + } + }) + } + } +} + +func TestPrometheusGridFallback(t *testing.T) { + b, err := NewClient("test", bo.New(), nil, nil, nil, nil) + if err != nil { + t.Fatal(err) + } + c := b.(*Client) + for _, tc := range []struct{ start, step string }{ + {"NaN", "15"}, + {"+Inf", "15"}, + {"1704067200.0001", "15"}, + {"1704067200", "0"}, + {"1704067200", "-1"}, + {"1704067200", "0.0001"}, + {"99999999999999", "1"}, + {"0001-01-01T00:00:00Z", "15"}, + } { + t.Run(fmt.Sprint(tc), func(t *testing.T) { + v := url.Values{"query": {"up"}, "start": {tc.start}, "end": {"1704067500"}, "step": {tc.step}} + r := httptest.NewRequest("GET", promPath+"/api/v1/query_range?"+v.Encode(), nil) + if _, _, _, err := c.ParseTimeRangeQuery(r); err == nil { + t.Fatal("unsupported grid was admitted") + } + }) + } + c.Configuration().DoesShard = true + r := httptest.NewRequest("GET", promPath+"/api/v1/query_range?query=up&start=1704067207&end=1704067507&step=15", nil) + if trq, _, _, err := c.ParseTimeRangeQuery(r); err != nil || trq.Phase != 0 { + t.Fatalf("step-aligned grid must support sharding: query=%+v err=%v", trq, err) + } +} + +func TestPrometheusCacheIdentity(t *testing.T) { + for _, method := range []string{"GET", "POST"} { + t.Run(method, func(t *testing.T) { + var calls atomic.Int32 + origin := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + id := calls.Add(1) + _ = r.ParseForm() + at, err := time.Parse(time.RFC3339Nano, r.Form.Get("start")) + if err != nil { + t.Error(err) + } + w.Header().Set("Content-Type", "application/json") + fmt.Fprintf(w, `{"status":"success","data":{"resultType":"matrix","result":[{"metric":{"job":"fixture"},"values":[[%d,"%d"]]}]}}`, at.Unix(), id) + }) + h := newHTTPHarness(t, origin) + h.resources.BackendOptions.FastForwardDisable = true + start := time.Now().UTC().Add(-5 * time.Minute).Truncate(15 * time.Second).Add(7 * time.Second) + for _, tc := range []struct { + name string + param, value, header string + }{ + {"default", "", "", ""}, + {"db_a", "db", "alpha", ""}, + {"db_b", "db", "beta", ""}, + {"header_a", "", "alpha", databaseHeader}, + {"header_b", "", "beta", databaseHeader}, + {"lookback_a", "lookback", "1m", ""}, + {"lookback_b", "lookback", "2m", ""}, + {"authorization_a", "", "Basic Zm9vOmJhcg==", "Authorization"}, + {"authorization_b", "", "Basic YmFyOmJheg==", "Authorization"}, + {"greptime_auth", "", "Basic YmFyOmJheg==", "X-Greptime-Auth"}, + } { + t.Run(tc.name, func(t *testing.T) { + v := url.Values{"query": {"up"}, "start": {start.Format(time.RFC3339Nano)}, "end": {start.Add(30 * time.Second).Format(time.RFC3339Nano)}, "step": {"15"}} + hdr := make(http.Header) + if tc.param != "" { + v.Set(tc.param, tc.value) + } + if tc.header != "" { + hdr.Set(tc.header, tc.value) + } + before := calls.Load() + var first string + for _, status := range []string{"kmiss", "hit"} { + w := h.promQuery(t, method, "query_range", v, hdr) + _, got := headers.ParseResultEngineStatus(w.Header().Get(headers.NameTricksterResult)) + if w.Code != 200 || got != status { + t.Fatalf("%d %s %s", w.Code, got, w.Body.String()) + } + if status == "kmiss" { + first = w.Body.String() + } else if first != w.Body.String() { + t.Fatal("cached response changed") + } + } + if calls.Load() != before+1 { + t.Fatal("cache identity collided or failed to reuse") + } + }) + } + }) + } +} + +func TestPrometheusUnsupportedPreservesRequest(t *testing.T) { + for _, tc := range []struct{ method, path, query, body, contentType string }{ + {"PUT", "query_range", "query=up", "original", "text/plain"}, + {"GET", "query_range", "query=up&limit=1", "", ""}, + {"GET", "query_range", "query=up&query=down", "", ""}, + {"POST", "query_range", "", "query=up", "application/json"}, + {"POST", "query_range", "", "query=%XX", "application/x-www-form-urlencoded"}, + {"POST", "label/job/values", "", "match[]=up", "application/x-www-form-urlencoded"}, + } { + t.Run(tc.method+tc.path+tc.query+tc.body, func(t *testing.T) { + origin := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + body, _ := io.ReadAll(r.Body) + if r.Method != tc.method || r.URL.RawQuery != tc.query || string(body) != tc.body { + t.Errorf("rewrote unsupported request: %s %s %q", r.Method, r.URL.RawQuery, body) + } + http.Error(w, "origin response", http.StatusBadRequest) + }) + h := newHTTPHarness(t, origin) + r := httptest.NewRequest(tc.method, promPath+"/api/v1/"+tc.path+"?"+tc.query, strings.NewReader(tc.body)) + r.Header.Set("Content-Type", tc.contentType) + res := h.resources + pc := h.client.DefaultPathConfigs(res.BackendOptions).Match(tc.method, r.URL.Path) + r = request.SetResources(r, request.NewResources(res.BackendOptions, pc, res.CacheConfig, res.CacheClient, h.client, res.Tracer)) + w := httptest.NewRecorder() + h.client.Handlers()[pc.HandlerName].ServeHTTP(w, r) + if w.Code != http.StatusBadRequest || w.Body.String() != "origin response\n" { + t.Fatalf("lost origin error: %d %s", w.Code, w.Body.String()) + } + }) + } +} + +func TestPrometheusMergePlanContract(t *testing.T) { + b, err := NewClient("test", bo.New(), nil, nil, nil, nil) + if err != nil { + t.Fatal(err) + } + c := b.(*Client) + t.Run("URL expression selects the plan", func(t *testing.T) { + r := httptest.NewRequest("POST", promPath+"/api/v1/query?query=avg(up)&db=public", strings.NewReader("query=min(up)&db=ignored")) + r.Header.Set("Content-Type", "application/x-www-form-urlencoded") + plan, err := c.PlanTSMMerge(r, "min(up)") + if err != nil { + t.Fatal(err) + } + if plan.OriginalQuery != "avg(up)" || plan.Reduction.Kind != merge.TSMReductionWeightedAverage { + t.Fatalf("wrong authoritative expression: %+v", plan) + } + body, _ := io.ReadAll(r.Body) + if string(body) != "query=min(up)&db=ignored" || r.URL.Query().Get("query") != "avg(up)" { + t.Fatal("planner changed caller request") + } + }) + for _, query := range []string{"group(up)", "topk(2, up)", "sort(group(up))"} { + t.Run(query, func(t *testing.T) { + r := httptest.NewRequest("GET", promPath+"/api/v1/query?query="+url.QueryEscape(query), nil) + plan, err := c.PlanTSMMerge(r, query) + if err != nil { + t.Fatal(err) + } + if !plan.StripInjectedLabels { + t.Fatal("aggregation retains per-backend routing labels") + } + }) + } + t.Run("quantile retains the origin metric name", func(t *testing.T) { + trq := ×eries.TimeRangeQuery{Statement: "up"} + ts, err := c.Modeler().WireUnmarshaler([]byte(`{"status":"success","data":{"resultType":"vector","result":[{"metric":{"__name__":"up","host":"a"},"value":[100,"1"]},{"metric":{"__name__":"up","host":"b"},"value":[100,"3"]},{"metric":{"__name__":"up","host":"c"},"value":[100,"10"]}]}}`), trq) + if err != nil { + t.Fatal(err) + } + c.FinalizeTSMMerge("quantile(0.5, up)", ts) + ds := ts.(*dataset.DataSet) + if len(ds.Results) != 1 || len(ds.Results[0].SeriesList) != 1 { + t.Fatal("wrong quantile result") + } + s := ds.Results[0].SeriesList[0] + if s.Header.Name != "up" || s.Header.Tags["__name__"] != "up" || s.Points[0].Values[0] != "3" { + t.Fatalf("lost Greptime metric identity: %+v", s) + } + }) +} + +func TestGreptimeMetricName(t *testing.T) { + for _, tc := range []struct { + query, name string + known bool + }{ + {"up", "up", true}, + {"sum(up)", "up", true}, + {"sum by () (up)", "", true}, + {"sum by (host) (up)", "", true}, + {"sum by (__name__) (up)", "up", true}, + {"sum without (host) (up)", "up", true}, + {"sum without (__name__) (up)", "", true}, + {"-up", "", true}, + {"up + 1", "", true}, + {"count(up) or vector(0)", "up", true}, + {"up and down", "up", true}, + {"up unless down", "up", true}, + {"rate(up[5m])", "up", true}, + {"max_over_time((up)[5m:])", "up", true}, + {`{"__name__"="quoted.metric"}`, "quoted.metric", true}, + {`{__name__=~"up|down"}`, "", false}, + {`{host="a"}`, "", false}, + {`quantile by (host) (0.5, {__name__=~"up|down"})`, "", false}, + {`quantile(0.5, {__name__=~"up|down"}) / 2`, "", false}, + {`histogram_fraction(0, 1, up)`, "up", true}, + {"sum(", "", false}, + {`label_replace(up, "a", "b", "c", "d")`, "up", true}, + } { + t.Run(tc.query, func(t *testing.T) { + name, known := greptimeMetricName(tc.query) + if name != tc.name || known != tc.known { + t.Fatalf("got %q %t, want %q %t", name, known, tc.name, tc.known) + } + }) + } +} + +func TestPrometheusCompressedError(t *testing.T) { + const body = `{"status":"error","errorType":"internal","error":"origin failed"}` + for _, encoding := range []string{"gzip", "br", "zstd", "deflate"} { + for _, method := range []string{"GET", "POST"} { + t.Run(encoding+method, func(t *testing.T) { + var encoded bytes.Buffer + init, header := ep.GetEncoderInitializer(encoding) + encoder := init(&encoded, 1) + if _, err := io.WriteString(encoder, body); err != nil { + t.Fatal(err) + } + if err := encoder.Close(); err != nil { + t.Fatal(err) + } + h := newHTTPHarness(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + w.Header().Set("Content-Encoding", header) + w.Header().Set("X-Origin-Error", "retained") + w.Header().Set("Content-Length", strconv.Itoa(encoded.Len())) + w.WriteHeader(http.StatusServiceUnavailable) + _, _ = w.Write(encoded.Bytes()) + })) + start := time.Now().Add(-5 * time.Minute).Truncate(15 * time.Second).Add(125 * time.Millisecond) + v := url.Values{"query": {"up"}, "start": {start.Format(time.RFC3339Nano)}, "end": {start.Add(30 * time.Second).Format(time.RFC3339Nano)}, "step": {"15"}} + w := h.promQuery(t, method, "query_range", v, http.Header{"Accept-Encoding": {"identity"}}) + if w.Code != http.StatusServiceUnavailable || w.Body.String() != body || w.Header().Get("Content-Encoding") != "" || w.Header().Get("X-Origin-Error") != "retained" { + t.Fatalf("compressed origin error corrupted: %d %v %q", w.Code, w.Header(), w.Body.String()) + } + }) + } + } +} diff --git a/pkg/backends/greptimedb/render.go b/pkg/backends/greptimedb/render.go new file mode 100644 index 000000000..b5ea8aad4 --- /dev/null +++ b/pkg/backends/greptimedb/render.go @@ -0,0 +1,93 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 greptimedb + +import ( + "errors" + "strings" + + "github.com/trickstercache/trickster/v2/pkg/parsing/sqlanalyzer/cockroach" + "github.com/trickstercache/trickster/v2/pkg/parsing/sqlscan" +) + +var errUnrenderable = errors.New("statement has no faithful GreptimeDB re-spelling") + +type textEdit struct { + start, end int + text string +} + +func postRender(rendered string) (string, error) { + masked := cockroach.MaskPlaceholders(rendered) + scanner := sqlscan.New(masked, sqlscan.Options{}) + var tokens []sqlscan.Token + for { + token, more := scanner.Next() + if !more { + break + } + tokens = append(tokens, token) + } + text := func(i int) string { + if i >= len(tokens) { + return "" + } + return masked[tokens[i].Start:tokens[i].End] + } + var edits []textEdit + for i, token := range tokens { + if token.Kind != sqlscan.Word { + continue + } + switch text(i) { + case "STRING": + edits = append(edits, textEdit{token.Start, token.End, "TEXT"}) + case "BYTES": + edits = append(edits, textEdit{token.Start, token.End, "BYTEA"}) + case "bucket": + // DataFusion reserves this word, but Cockroach drops its identifier quotes. + edits = append(edits, textEdit{token.Start, token.End, `"bucket"`}) + case "extract": + if text(i+1) != "(" || i+3 >= len(tokens) || tokens[i+2].Kind != sqlscan.String || text(i+3) != "," { + continue + } + field := strings.Trim(text(i+2), "'") + if field == "" || len(field)+2 != tokens[i+2].End-tokens[i+2].Start { + return "", errUnrenderable + } + for _, char := range field { + if char != '_' && (char < 'a' || char > 'z') && (char < 'A' || char > 'Z') { + return "", errUnrenderable + } + } + edits = append(edits, textEdit{tokens[i+2].Start, tokens[i+3].End, field + " FROM"}) + } + } + if len(edits) == 0 { + return rendered, nil + } + var out strings.Builder + out.Grow(len(rendered) + 8*len(edits)) + at := 0 + for _, edit := range edits { + out.WriteString(rendered[at:edit.start]) + out.WriteString(edit.text) + at = edit.end + } + out.WriteString(rendered[at:]) + return out.String(), nil +} diff --git a/pkg/backends/greptimedb/routes.go b/pkg/backends/greptimedb/routes.go new file mode 100644 index 000000000..e97b93ca7 --- /dev/null +++ b/pkg/backends/greptimedb/routes.go @@ -0,0 +1,88 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 greptimedb + +import ( + "net/http" + + bo "github.com/trickstercache/trickster/v2/pkg/backends/options" + "github.com/trickstercache/trickster/v2/pkg/backends/prometheus" + "github.com/trickstercache/trickster/v2/pkg/backends/providers" + "github.com/trickstercache/trickster/v2/pkg/proxy/engines" + "github.com/trickstercache/trickster/v2/pkg/proxy/handlers" + "github.com/trickstercache/trickster/v2/pkg/proxy/headers" + "github.com/trickstercache/trickster/v2/pkg/proxy/methods" + "github.com/trickstercache/trickster/v2/pkg/proxy/paths/matching" + po "github.com/trickstercache/trickster/v2/pkg/proxy/paths/options" + "github.com/trickstercache/trickster/v2/pkg/proxy/urls" +) + +func (c *Client) RegisterHandlers(handlers.Lookup) { + lookup := c.Client.HandlerLookup() + lookup["sql"] = http.HandlerFunc(c.QueryHandler) + lookup["health"] = http.HandlerFunc(c.HealthHandler) + lookup["proxy"] = http.HandlerFunc(c.ProxyHandler) + c.TimeseriesBackend.RegisterHandlers(lookup) +} + +// DefaultPathConfigs preserves GreptimeDB's paths, including non-query APIs. +func (c *Client) DefaultPathConfigs(o *bo.Options) po.List { + paths := po.List{{ + Path: "/", HandlerName: providers.Proxy, Methods: methods.AllHTTPMethods(), + MatchType: matching.PathMatchTypePrefix, MatchTypeName: matching.PathMatchNamePrefix, + }, { + // The exact route must not mask catch-all passthrough for other methods. + Path: "/v1/sql", HandlerName: "sql", Methods: methods.AllHTTPMethods(), + MatchType: matching.PathMatchTypeExact, MatchTypeName: matching.PathMatchNameExact, + CacheKeyParams: []string{"sql", "db", "format", "epoch", "limit", "compression"}, + CacheKeyHeaders: []string{"X-Greptime-Db-Name", "X-Greptime-Timezone", "X-Greptime-Auth"}, + }} + hooks := promHooks() + promPaths := prometheus.Without(prometheus.SupportedPaths(o), "alerts", "proxycache", "admin", "proxy") + promPaths = prometheus.WithPathPrefix(promPaths, hooks.PathPrefix) + promPaths = prometheus.WithCacheKeyParams(promPaths, hooks.CacheKeyParams...) + promPaths = prometheus.WithCacheKeyHeaders(promPaths, hooks.CacheKeyHeaders...) + for _, p := range promPaths { + // Only the origin may authorize sharing authenticated responses. + delete(p.ResponseHeaders, headers.NameCacheControl) + // The request hook relays unsupported methods instead of masking the catch-all. + p.Methods = methods.AllHTTPMethods() + if o != nil && p.HandlerName == "query" { + o.FastForwardPath = p.Clone() + } + } + return append(paths, promPaths...) +} + +// MergeablePaths excludes unsupported endpoints and the separate SQL surface. +func (c *Client) MergeablePaths() []string { + paths := c.Client.MergeablePaths() + supported := c.DefaultPathConfigs(nil) + out := make([]string, 0, len(paths)) + for _, path := range paths { + if p := supported.Match(http.MethodGet, path); p != nil && p.HandlerName != providers.Proxy { + out = append(out, path) + } + } + return out +} + +// ProxyHandler relays HTTP requests without caching. +func (c *Client) ProxyHandler(w http.ResponseWriter, r *http.Request) { + r.URL = urls.BuildUpstreamURL(r, c.BaseUpstreamURL()) + engines.DoProxy(w, r, true) +} diff --git a/pkg/backends/greptimedb/session.go b/pkg/backends/greptimedb/session.go new file mode 100644 index 000000000..fb2164148 --- /dev/null +++ b/pkg/backends/greptimedb/session.go @@ -0,0 +1,49 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 greptimedb + +import "github.com/trickstercache/trickster/v2/pkg/proxy/pgwire" + +var sessionSettings = pgwire.SessionSettings{ + Tracked: map[string]struct{}{ + "timezone": {}, "datestyle": {}, "intervalstyle": {}, "bytea_output": {}, "search_path": {}, + }, + Neutral: map[string]struct{}{ + "application_name": {}, "statement_timeout": {}, "client_encoding": {}, + // Accepted as no-ops; float output and string parsing are not configurable this way. + "extra_float_digits": {}, "standard_conforming_strings": {}, + }, + Aliases: map[string]string{"time_zone": "timezone"}, + LocalPersists: true, UnconfirmedStartup: true, +} + +func (engine) SessionDefaultsProbe() pgwire.SessionDefaultsProbe { + return pgwire.SessionDefaultsProbe{ + SQL: "SHOW TIMEZONE; SHOW DateStyle; SHOW IntervalStyle", + Names: []string{"timezone", "datestyle", "intervalstyle"}, + } +} + +func (engine) SessionSettings() pgwire.SessionSettings { return sessionSettings } + +func (engine) TimeSemantics() pgwire.TimeSemantics { + return pgwire.TimeSemantics{ + NaiveTimestampsAreUTC: true, LosslessFloatText: true, + AssumedDateStyle: "ISO, MDY", AssumedIntervalStyle: "postgres", + AssumedStandardConformingStrings: "on", AssumedIntegerDatetimes: "on", + } +} diff --git a/pkg/backends/greptimedb/sql/request.go b/pkg/backends/greptimedb/sql/request.go new file mode 100644 index 000000000..044848562 --- /dev/null +++ b/pkg/backends/greptimedb/sql/request.go @@ -0,0 +1,189 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 sql adapts GreptimeDB HTTP SQL requests to the shared SQL analyzer. +package sql + +import ( + "errors" + "maps" + "mime" + "net/http" + "net/url" + "strings" + "time" + + "github.com/trickstercache/trickster/v2/pkg/parsing/sqlanalyzer" + "github.com/trickstercache/trickster/v2/pkg/parsing/sqlscan" + "github.com/trickstercache/trickster/v2/pkg/proxy/request" + "github.com/trickstercache/trickster/v2/pkg/proxy/urls" + "github.com/trickstercache/trickster/v2/pkg/timeseries" +) + +var ( + errRequest = errors.New("unsupported GreptimeDB SQL request") + errObject = errors.New("GreptimeDB SQL request requires object caching") +) + +type sqlRequest struct { + query, form url.Values + values url.Values + body []byte +} + +func extract(r *http.Request) (*sqlRequest, error) { + if r == nil || r.URL == nil || (r.Method != http.MethodGet && r.Method != http.MethodPost) { + return nil, errRequest + } + query, err := url.ParseQuery(r.URL.RawQuery) + if err != nil { + return nil, err + } + out := &sqlRequest{query: query, values: maps.Clone(query)} + if r.Method == http.MethodPost { + media, _, err := mime.ParseMediaType(r.Header.Get("Content-Type")) + if err != nil || media != "application/x-www-form-urlencoded" { + return nil, errRequest + } + out.body, err = request.GetBody(r) + if err != nil { + return nil, err + } + out.form, err = url.ParseQuery(string(out.body)) + if err != nil { + return nil, err + } + for k, v := range out.form { + // An explicitly empty URL value still overrides the form value. + if _, exists := query[k]; !exists { + out.values[k] = v + } + } + } + for _, values := range []url.Values{out.query, out.form} { + for _, v := range values { + if len(v) != 1 { + return nil, errRequest + } + } + } + if out.values.Get("sql") == "" { + return nil, errRequest + } + return out, nil +} + +// ParseTimeRangeQuery uses the same analyzer as the provider's native SQL path. +func ParseTimeRangeQuery(r *http.Request, analyzer sqlanalyzer.DialectAnalyzer, +) (*timeseries.TimeRangeQuery, *timeseries.RequestOptions, bool, error) { + input, err := extract(r) + if err != nil || analyzer == nil { + return nil, nil, false, errRequest + } + now := time.Now() + statement := input.values.Get("sql") + if !singleSelect(statement) { + return nil, nil, false, errRequest + } + analysis := analyzer.Analyze(statement, now) + if analysis.Mode == sqlanalyzer.CacheModeNone { + return nil, nil, false, errRequest + } + trq := sqlanalyzer.NewTimeRangeQuery(statement) + trq.OriginalBody = input.body + trq.CacheKeyElements["greptime.http"] = input.values.Encode() + ro := ×eries.RequestOptions{FastForwardDisable: true, ResponseContentType: "application/json", FallbackToProxyOnError: true} + format := strings.ToLower(input.values.Get("format")) + _, limited := input.values["limit"] + known := true + for key := range input.values { + switch key { + case "sql", "db", "format", "epoch", "compression": + default: + known = false + } + } + if analysis.Mode != sqlanalyzer.CacheModeDelta || analysis.Plan == nil || limited || !known || + (format != "" && format != "greptimedb_v1") { + return trq, ro, true, errObject + } + plan := analysis.Plan + plan.ApplyToQuery(trq) + trq.Extent = plan.RequestExtent(now) + input.values.Set("sql", plan.CanonicalSQL) + trq.CacheKeyElements["greptime.http"] = input.values.Encode() + trq.CacheKeyElements["sql"] = plan.CanonicalSQL + trq.TemplateURL = urls.Clone(r.URL) + trq.TemplateURL.RawQuery = input.values.Encode() + ro.BaseTimestampFieldName = plan.TimeColumn + if trq.BackfillTolerance == 0 { + bf := time.Minute + if res := request.GetResources(r); res != nil && res.BackendOptions != nil { + bf = time.Duration(res.BackendOptions.BackfillTolerance) + } + if plan.UpperBound == nil && bf < trq.Step { + bf = trq.Step + } + trq.BackfillTolerance = bf + } + return trq, ro, true, nil +} + +func singleSelect(statement string) bool { + scanner := sqlscan.New(statement, sqlscan.Options{}) + seen, ended := false, false + for { + token, more := scanner.Next() + if !more { + return seen && !scanner.Unterminated + } + if token.Kind == sqlscan.Punct && scanner.Text(token) == ";" { + ended = true + continue + } + if ended || (!seen && !scanner.IsWord(token, "select")) { + return false + } + seen = true + } +} + +// SetExtent rewrites the effective SQL field, leaving all other parameters intact. +func SetExtent(r *http.Request, trq *timeseries.TimeRangeQuery, extent *timeseries.Extent) error { + if trq == nil || extent == nil { + return errRequest + } + plan, ok := trq.ParsedQuery.(*sqlanalyzer.QueryPlan) + if !ok || plan == nil { + return errRequest + } + input, err := extract(r) + if err != nil { + return err + } + statement, err := plan.RenderExtent(*extent) + if err != nil { + return err + } + if _, inURL := input.query["sql"]; inURL { + input.query.Set("sql", statement) + r.URL.RawQuery = input.query.Encode() + } else { + input.form.Set("sql", statement) + request.SetBody(r, []byte(input.form.Encode())) + } + return nil +} diff --git a/pkg/backends/greptimedb/sql/request_test.go b/pkg/backends/greptimedb/sql/request_test.go new file mode 100644 index 000000000..d6e120fab --- /dev/null +++ b/pkg/backends/greptimedb/sql/request_test.go @@ -0,0 +1,44 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 sql + +import "testing" + +func TestSingleSelect(t *testing.T) { + for _, test := range []struct { + sql string + want bool + }{ + {"SELECT 1", true}, + {"/* lead ; */ SELECT ';' AS x; -- tail ;", true}, + {"SELECT $$;$$;", true}, + {`SELECT E'\';';`, true}, + {"SELECT 1; SELECT 2", false}, + {"SELECT 1; INSERT INTO t VALUES (2)", false}, + {"INSERT INTO t VALUES (2)", false}, + {"WITH t AS (DELETE FROM m RETURNING *) SELECT * FROM t", false}, + {"", false}, + {"SELECT 'unterminated", false}, + {"SELECT 1 /*", false}, + } { + t.Run(test.sql, func(t *testing.T) { + if got := singleSelect(test.sql); got != test.want { + t.Fatalf("got %v want %v", got, test.want) + } + }) + } +} diff --git a/pkg/backends/greptimedb/testdata/compatibility/README.md b/pkg/backends/greptimedb/testdata/compatibility/README.md new file mode 100644 index 000000000..e57ce03eb --- /dev/null +++ b/pkg/backends/greptimedb/testdata/compatibility/README.md @@ -0,0 +1,25 @@ +# GreptimeDB Compatibility Corpus + +`v1.json` is run by `pkg/testutil/sqlcompat`, the same schema, macro coverage, +plan assertions and render/read-back checks used for PostgreSQL. Do not change +a released corpus's expectations without versioning the incompatible change. + +The eight Grafana cases were captured from `executedQueryString` using the +bundled PostgreSQL datasource in Grafana 13.1.3 against GreptimeDB's official +`nightly-20260923-e91faa9df` image. `timescaledb` was false, the query interval +was five minutes, and the fixed request window was +`2026-09-19T00:00:03.123Z` through `2026-09-19T03:00:07.456Z`. Each query succeeded +directly. `macro_source` records the submitted SQL, not a reconstructed macro. +All fourteen PostgreSQL macro families are represented. Minimum dashboard +interval remains one minute. + +The hand-written cases cover DataFusion's interval and compact `date_bin`, +`date_part`, calendar/submicrosecond fallbacks, session timezone, `RANGE/ALIGN`, +TQL, volatility and writes. A delta case must keep its canonical identity and +read back exactly the inclusive cache extent requested by the renderer. + +These are pgwire analyzer expectations. HTTP SQL additionally refuses delta +rewrites that drop partial buckets. MySQL uses a different grammar and its +own analyzer and live typed-result tests. Extended pgwire queries are relayed, +not cached; Builder-mode duplicate captures and advanced clause rewriting +are explicitly omitted. diff --git a/pkg/backends/greptimedb/testdata/compatibility/v1.json b/pkg/backends/greptimedb/testdata/compatibility/v1.json new file mode 100644 index 000000000..c578d86ff --- /dev/null +++ b/pkg/backends/greptimedb/testdata/compatibility/v1.json @@ -0,0 +1,417 @@ +{ + "schema_version": 1, + "corpus_version": "greptimedb-grafana-v1", + "minimum_interval": "1m", + "omissions": [ + "builder-mode-captures", + "extended-protocol-statements", + "RANGE ALIGN delta caching", + "TQL delta caching" + ], + "cases": [ + { + "name": "grafana_time", + "macro_source": [ + "SELECT $__time(pickup_datetime), total_amount FROM trips WHERE $__timeFilter(pickup_datetime) ORDER BY 1" + ], + "query_origin": "Grafana PostgreSQL datasource executedQueryString, direct GreptimeDB official nightly-20260923-e91faa9df, timescaledb=false", + "grafana_version": "13.1.3", + "session_time_zone": "UTC", + "expanded_sql": "SELECT pickup_datetime AS \"time\", total_amount FROM trips WHERE pickup_datetime BETWEEN '2026-09-19T00:00:03.123Z' AND '2026-09-19T03:00:07.456Z' ORDER BY 1", + "expected": { + "cache_mode": "object", + "analysis_reason": "unsupported_bucket", + "canonical_policy": "none", + "extent_rendering": false + }, + "rationale": "The projection supplies no reusable bucket cadence; cache only the complete result." + }, + { + "name": "grafana_time_epoch", + "macro_source": [ + "SELECT $__timeEpoch(pickup_datetime), total_amount FROM trips WHERE $__timeFilter(pickup_datetime) ORDER BY 1" + ], + "query_origin": "Grafana PostgreSQL datasource executedQueryString, direct GreptimeDB official nightly-20260923-e91faa9df, timescaledb=false", + "grafana_version": "13.1.3", + "session_time_zone": "UTC", + "expanded_sql": "SELECT extract(epoch from pickup_datetime) as \"time\", total_amount FROM trips WHERE pickup_datetime BETWEEN '2026-09-19T00:00:03.123Z' AND '2026-09-19T03:00:07.456Z' ORDER BY 1", + "expected": { + "cache_mode": "object", + "analysis_reason": "unsupported_bucket", + "canonical_policy": "none", + "extent_rendering": false + }, + "rationale": "The projection supplies no reusable bucket cadence; cache only the complete result." + }, + { + "name": "grafana_time_group", + "macro_source": [ + "SELECT $__timeGroup(pickup_datetime,'5m') AS time, count(*) AS trips FROM trips WHERE $__timeFilter(pickup_datetime) GROUP BY 1 ORDER BY 1" + ], + "query_origin": "Grafana PostgreSQL datasource executedQueryString, direct GreptimeDB official nightly-20260923-e91faa9df, timescaledb=false", + "grafana_version": "13.1.3", + "session_time_zone": "UTC", + "expanded_sql": "SELECT floor(extract(epoch from pickup_datetime)/300)*300 AS time, count(*) AS trips FROM trips WHERE pickup_datetime BETWEEN '2026-09-19T00:00:03.123Z' AND '2026-09-19T03:00:07.456Z' GROUP BY 1 ORDER BY 1", + "expected": { + "cache_mode": "delta", + "analysis_reason": "delta_cacheable", + "cadence": "5m", + "phase": "0s", + "input_unit": "rfc3339", + "output_unit": "unix_seconds", + "lower_bound": "2026-09-19T00:05:00Z", + "lower_inclusive": true, + "upper_bound": "2026-09-19T03:00:00Z", + "upper_inclusive": false, + "output_column": "time", + "group_columns": [], + "canonical_policy": "range_independent", + "extent_rendering": true, + "canonical_contains": [ + "<$TS1$>", + "<$TS2$>" + ] + }, + "rationale": "Fixed five-minute buckets normalize the deliberately unaligned range to complete buckets." + }, + { + "name": "grafana_time_group_alias", + "macro_source": [ + "SELECT $__timeGroupAlias(pickup_datetime,$__interval), count(*) AS trips FROM trips WHERE pickup_datetime >= $__timeFrom() AND pickup_datetime < $__timeTo() GROUP BY 1 ORDER BY 1" + ], + "query_origin": "Grafana PostgreSQL datasource executedQueryString, direct GreptimeDB official nightly-20260923-e91faa9df, timescaledb=false", + "grafana_version": "13.1.3", + "session_time_zone": "UTC", + "expanded_sql": "SELECT floor(extract(epoch from pickup_datetime)/300)*300 AS \"time\", count(*) AS trips FROM trips WHERE pickup_datetime >= '2026-09-19T00:00:03.123Z' AND pickup_datetime < '2026-09-19T03:00:07.456Z' GROUP BY 1 ORDER BY 1", + "expected": { + "cache_mode": "delta", + "analysis_reason": "delta_cacheable", + "cadence": "5m", + "phase": "0s", + "input_unit": "rfc3339", + "output_unit": "unix_seconds", + "lower_bound": "2026-09-19T00:05:00Z", + "lower_inclusive": true, + "upper_bound": "2026-09-19T03:00:00Z", + "upper_inclusive": false, + "output_column": "time", + "group_columns": [], + "canonical_policy": "range_independent", + "extent_rendering": true, + "canonical_contains": [ + "<$TS1$>", + "<$TS2$>" + ] + }, + "rationale": "Fixed five-minute buckets normalize the deliberately unaligned range to complete buckets." + }, + { + "name": "grafana_unix_filter", + "macro_source": [ + "SELECT pickup_epoch AS time, total_amount FROM trips WHERE $__unixEpochFilter(pickup_epoch) ORDER BY 1" + ], + "query_origin": "Grafana PostgreSQL datasource executedQueryString, direct GreptimeDB official nightly-20260923-e91faa9df, timescaledb=false", + "grafana_version": "13.1.3", + "session_time_zone": "UTC", + "expanded_sql": "SELECT pickup_epoch AS time, total_amount FROM trips WHERE pickup_epoch >= 1789776003 AND pickup_epoch <= 1789786807 ORDER BY 1", + "expected": { + "cache_mode": "object", + "analysis_reason": "unsupported_bucket", + "canonical_policy": "none", + "extent_rendering": false + }, + "rationale": "The projection supplies no reusable bucket cadence; cache only the complete result." + }, + { + "name": "grafana_unix_nano_filter", + "macro_source": [ + "SELECT pickup_epoch AS time, total_amount FROM trips WHERE $__unixEpochNanoFilter(pickup_epoch * 1000000000) ORDER BY 1" + ], + "query_origin": "Grafana PostgreSQL datasource executedQueryString, direct GreptimeDB official nightly-20260923-e91faa9df, timescaledb=false", + "grafana_version": "13.1.3", + "session_time_zone": "UTC", + "expanded_sql": "SELECT pickup_epoch AS time, total_amount FROM trips WHERE pickup_epoch * 1000000000 >= 1789776003123000000 AND pickup_epoch * 1000000000 <= 1789786807456000000 ORDER BY 1", + "expected": { + "cache_mode": "object", + "analysis_reason": "unsupported_bucket", + "canonical_policy": "none", + "extent_rendering": false + }, + "rationale": "The projection supplies no reusable bucket cadence; cache only the complete result." + }, + { + "name": "grafana_unix_group", + "macro_source": [ + "SELECT $__unixEpochGroup(pickup_epoch,'5m') AS time, count(*) AS trips FROM trips WHERE pickup_epoch >= $__unixEpochFrom() AND pickup_epoch < $__unixEpochTo() GROUP BY 1 ORDER BY 1" + ], + "query_origin": "Grafana PostgreSQL datasource executedQueryString, direct GreptimeDB official nightly-20260923-e91faa9df, timescaledb=false", + "grafana_version": "13.1.3", + "session_time_zone": "UTC", + "expanded_sql": "SELECT floor((pickup_epoch)/300)*300 AS time, count(*) AS trips FROM trips WHERE pickup_epoch >= 1789776003 AND pickup_epoch < 1789786807 GROUP BY 1 ORDER BY 1", + "expected": { + "cache_mode": "delta", + "analysis_reason": "delta_cacheable", + "cadence": "5m", + "phase": "0s", + "input_unit": "unix_seconds", + "output_unit": "unix_seconds", + "lower_bound": "2026-09-19T00:05:00Z", + "lower_inclusive": true, + "upper_bound": "2026-09-19T03:00:00Z", + "upper_inclusive": false, + "output_column": "time", + "group_columns": [], + "canonical_policy": "range_independent", + "extent_rendering": true, + "canonical_contains": [ + "<$TS1$>", + "<$TS2$>" + ] + }, + "rationale": "Fixed five-minute buckets normalize the deliberately unaligned range to complete buckets." + }, + { + "name": "grafana_unix_group_alias", + "macro_source": [ + "SELECT $__unixEpochGroupAlias(pickup_epoch,'5m'), count(*) AS trips FROM trips WHERE $__unixEpochFilter(pickup_epoch) GROUP BY 1 ORDER BY 1" + ], + "query_origin": "Grafana PostgreSQL datasource executedQueryString, direct GreptimeDB official nightly-20260923-e91faa9df, timescaledb=false", + "grafana_version": "13.1.3", + "session_time_zone": "UTC", + "expanded_sql": "SELECT floor((pickup_epoch)/300)*300 AS \"time\", count(*) AS trips FROM trips WHERE pickup_epoch >= 1789776003 AND pickup_epoch <= 1789786807 GROUP BY 1 ORDER BY 1", + "expected": { + "cache_mode": "delta", + "analysis_reason": "delta_cacheable", + "cadence": "5m", + "phase": "0s", + "input_unit": "unix_seconds", + "output_unit": "unix_seconds", + "lower_bound": "2026-09-19T00:05:00Z", + "lower_inclusive": true, + "upper_bound": "2026-09-19T03:00:00Z", + "upper_inclusive": false, + "output_column": "time", + "group_columns": [], + "canonical_policy": "range_independent", + "extent_rendering": true, + "canonical_contains": [ + "<$TS1$>", + "<$TS2$>" + ] + }, + "rationale": "Fixed five-minute buckets normalize the deliberately unaligned range to complete buckets." + }, + { + "name": "date_bin_interval", + "query_origin": "Hand-written DataFusion SQL regression", + "session_time_zone": "UTC", + "expanded_sql": "SELECT date_bin(INTERVAL '5 minutes',pickup_datetime) AS time,count(*) AS trips FROM trips WHERE pickup_datetime >= '2026-01-01T00:00:00Z' AND pickup_datetime < '2026-01-01T01:00:00Z' GROUP BY 1 ORDER BY 1", + "expected": { + "cache_mode": "delta", + "analysis_reason": "delta_cacheable", + "cadence": "5m", + "phase": "0s", + "input_unit": "rfc3339", + "output_unit": "timestamp", + "lower_bound": "2026-01-01T00:00:00Z", + "lower_inclusive": true, + "upper_bound": "2026-01-01T01:00:00Z", + "upper_inclusive": false, + "output_column": "time", + "group_columns": [], + "canonical_policy": "range_independent", + "extent_rendering": true, + "canonical_contains": [ + "<$TS1$>", + "<$TS2$>" + ] + }, + "rationale": "Fixed-width buckets render and read back with the original identity and exact requested extent." + }, + { + "name": "date_bin_compact", + "query_origin": "Hand-written DataFusion SQL regression", + "session_time_zone": "UTC", + "expanded_sql": "SELECT date_bin('5m',pickup_datetime) AS time,count(*) AS trips FROM trips WHERE pickup_datetime >= '2026-01-01T00:00:00Z' AND pickup_datetime < '2026-01-01T01:00:00Z' GROUP BY 1 ORDER BY 1", + "expected": { + "cache_mode": "delta", + "analysis_reason": "delta_cacheable", + "cadence": "5m", + "phase": "0s", + "input_unit": "rfc3339", + "output_unit": "timestamp", + "lower_bound": "2026-01-01T00:00:00Z", + "lower_inclusive": true, + "upper_bound": "2026-01-01T01:00:00Z", + "upper_inclusive": false, + "output_column": "time", + "group_columns": [], + "canonical_policy": "range_independent", + "extent_rendering": true, + "canonical_contains": [ + "<$TS1$>", + "<$TS2$>" + ] + }, + "rationale": "Fixed-width buckets render and read back with the original identity and exact requested extent." + }, + { + "name": "date_bin_interval_cast", + "query_origin": "Hand-written DataFusion SQL regression", + "session_time_zone": "UTC", + "expanded_sql": "SELECT date_bin('5 minutes'::interval,pickup_datetime) AS time,count(*) AS trips FROM trips WHERE pickup_datetime >= '2026-01-01T00:00:00Z' AND pickup_datetime < '2026-01-01T01:00:00Z' GROUP BY 1 ORDER BY 1", + "expected": { + "cache_mode": "delta", + "analysis_reason": "delta_cacheable", + "cadence": "5m", + "phase": "0s", + "input_unit": "rfc3339", + "output_unit": "timestamp", + "lower_bound": "2026-01-01T00:00:00Z", + "lower_inclusive": true, + "upper_bound": "2026-01-01T01:00:00Z", + "upper_inclusive": false, + "output_column": "time", + "group_columns": [], + "canonical_policy": "range_independent", + "extent_rendering": true, + "canonical_contains": [ + "<$TS1$>", + "<$TS2$>" + ] + }, + "rationale": "Fixed-width buckets render and read back with the original identity and exact requested extent." + }, + { + "name": "date_part_epoch", + "query_origin": "Hand-written DataFusion SQL regression", + "session_time_zone": "UTC", + "expanded_sql": "SELECT floor(date_part('epoch',pickup_datetime)/300)*300 AS time,count(*) AS trips FROM trips WHERE pickup_datetime >= '2026-01-01T00:00:00Z' AND pickup_datetime < '2026-01-01T01:00:00Z' GROUP BY 1 ORDER BY 1", + "expected": { + "cache_mode": "delta", + "analysis_reason": "delta_cacheable", + "cadence": "5m", + "phase": "0s", + "input_unit": "rfc3339", + "output_unit": "unix_seconds", + "lower_bound": "2026-01-01T00:00:00Z", + "lower_inclusive": true, + "upper_bound": "2026-01-01T01:00:00Z", + "upper_inclusive": false, + "output_column": "time", + "group_columns": [], + "canonical_policy": "range_independent", + "extent_rendering": true, + "canonical_contains": [ + "<$TS1$>", + "<$TS2$>" + ] + }, + "rationale": "Fixed-width buckets render and read back with the original identity and exact requested extent." + }, + { + "name": "calendar_month", + "query_origin": "Hand-written DataFusion SQL regression", + "session_time_zone": "UTC", + "expanded_sql": "SELECT date_trunc('month',pickup_datetime) AS time,count(*) AS trips FROM trips WHERE pickup_datetime >= '2026-01-01T00:00:00Z' AND pickup_datetime < '2026-01-01T01:00:00Z' GROUP BY 1 ORDER BY 1", + "expected": { + "cache_mode": "object", + "analysis_reason": "unsupported_bucket", + "canonical_policy": "none", + "extent_rendering": false + }, + "rationale": "Calendar months have no fixed cache cadence." + }, + { + "name": "fractional_compact", + "query_origin": "Hand-written DataFusion SQL regression", + "session_time_zone": "UTC", + "expanded_sql": "SELECT date_bin('0.5s',pickup_datetime) AS time,count(*) AS trips FROM trips WHERE pickup_datetime >= '2026-01-01T00:00:00Z' AND pickup_datetime < '2026-01-01T01:00:00Z' GROUP BY 1 ORDER BY 1", + "expected": { + "cache_mode": "object", + "analysis_reason": "unsupported_bucket", + "canonical_policy": "none", + "extent_rendering": false + }, + "rationale": "Compact widths require positive integer components." + }, + { + "name": "submicrosecond", + "query_origin": "Hand-written DataFusion SQL regression", + "session_time_zone": "UTC", + "expanded_sql": "SELECT date_bin('5ns',pickup_datetime) AS time,count(*) AS trips FROM trips WHERE pickup_datetime >= '2026-01-01T00:00:00Z' AND pickup_datetime < '2026-01-01T01:00:00Z' GROUP BY 1 ORDER BY 1", + "expected": { + "cache_mode": "object", + "analysis_reason": "unsupported_bucket", + "canonical_policy": "none", + "extent_rendering": false + }, + "rationale": "PostgreSQL text timestamps do not carry nanosecond buckets losslessly." + }, + { + "name": "unknown_zone", + "query_origin": "Hand-written DataFusion SQL regression", + "session_time_zone": "America/New_York", + "expanded_sql": "SELECT date_trunc('minute',pickup_datetime) AS time,count(*) AS trips FROM trips WHERE pickup_datetime >= '2026-01-01T00:00:00Z' AND pickup_datetime < '2026-01-01T01:00:00Z' GROUP BY 1 ORDER BY 1", + "expected": { + "cache_mode": "object", + "analysis_reason": "unsupported_bucket", + "canonical_policy": "none", + "extent_rendering": false + }, + "rationale": "date_trunc requires a verified UTC session." + }, + { + "name": "range_align", + "query_origin": "Hand-written DataFusion SQL regression", + "session_time_zone": "UTC", + "expanded_sql": "SELECT sum(fare_amount) RANGE '5m' FROM trips ALIGN '5m'", + "expected": { + "cache_mode": "object", + "analysis_reason": "invalid_sql", + "canonical_policy": "none", + "extent_rendering": false + }, + "rationale": "The dialect extension is only eligible for whole-result caching." + }, + { + "name": "range_volatile", + "query_origin": "Hand-written DataFusion SQL regression", + "session_time_zone": "UTC", + "expanded_sql": "SELECT random() RANGE '5m' FROM trips ALIGN '5m'", + "expected": { + "cache_mode": "none", + "analysis_reason": "nondeterministic", + "canonical_policy": "none", + "extent_rendering": false + }, + "rationale": "Lexical safety checks still apply when the AST parser cannot read a dialect clause." + }, + { + "name": "tql", + "query_origin": "Hand-written DataFusion SQL regression", + "session_time_zone": "UTC", + "expanded_sql": "TQL EVAL (0, 1, '1s') up", + "expected": { + "cache_mode": "none", + "analysis_reason": "invalid_sql", + "canonical_policy": "none", + "extent_rendering": false + }, + "rationale": "TQL is relayed without SQL cache rewrites." + }, + { + "name": "write", + "query_origin": "Hand-written DataFusion SQL regression", + "session_time_zone": "UTC", + "expanded_sql": "INSERT INTO trips VALUES (1)", + "expected": { + "cache_mode": "none", + "analysis_reason": "unsupported_statement", + "canonical_policy": "none", + "extent_rendering": false + }, + "rationale": "Mutations never enter a read-result cache." + } + ] +} diff --git a/pkg/backends/greptimedb/tsm.go b/pkg/backends/greptimedb/tsm.go new file mode 100644 index 000000000..6b92efb67 --- /dev/null +++ b/pkg/backends/greptimedb/tsm.go @@ -0,0 +1,188 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 greptimedb + +import ( + "errors" + "net/http" + "slices" + + "github.com/trickstercache/trickster/v2/pkg/backends" + "github.com/trickstercache/trickster/v2/pkg/backends/prometheus/promql" + "github.com/trickstercache/trickster/v2/pkg/proxy/params" + "github.com/trickstercache/trickster/v2/pkg/proxy/request" + "github.com/trickstercache/trickster/v2/pkg/timeseries" + "github.com/trickstercache/trickster/v2/pkg/timeseries/dataset" + "github.com/trickstercache/trickster/v2/pkg/timeseries/merge" + + "github.com/prometheus/prometheus/model/labels" + "github.com/prometheus/prometheus/promql/parser" +) + +var ( + _ backends.TSMMergeProvider = (*Client)(nil) + metricParser = parser.NewParser(parser.Options{EnableExperimentalFunctions: true}) +) + +// PlanTSMMerge normalizes Greptime's URL-first semantics before the shared +// planner chooses or rewrites any expression. Member handlers run too late. +func (c *Client) PlanTSMMerge(r *http.Request, _ string) (*merge.TSMMergePlan, error) { + req, err := request.Clone(r) + if err != nil { + return nil, err + } + if req == nil || !preparePromRequest(req) { + return nil, errors.New("unsupported GreptimeDB PromQL merge request") + } + values, _, _ := params.GetRequestValues(req) + query := values.Get("query") + plan, err := c.Client.PlanTSMMerge(req, query) + if err != nil { + return nil, err + } + if expr, err := promql.Parse(query); err == nil && expr.ContainsAggregation() { + plan.StripInjectedLabels = true + } + if plan.Finalizer.Enabled { + if _, known := greptimeMetricName(query); !known { + return nil, errors.New("cannot finalize GreptimeDB metric-name discovery across shards") + } + } + return plan, nil +} + +// FinalizeTSMMerge reuses the shared numeric reducers while preserving the +// metric-name contract of Greptime's HTTP response builder. +func (c *Client) FinalizeTSMMerge(query string, ts timeseries.Timeseries) { + c.Client.FinalizeTSMMerge(query, ts) + ds, ok := ts.(*dataset.DataSet) + if !ok || ds == nil { + return + } + name, known := greptimeMetricName(query) + if !known { + return + } + ds.UpdateLock.Lock() + defer ds.UpdateLock.Unlock() + for _, result := range ds.Results { + if result == nil { + continue + } + for _, series := range result.SeriesList { + if series == nil { + continue + } + if name == "" { + delete(series.Header.Tags, promql.MetricNameLabel) + } else { + if series.Header.Tags == nil { + series.Header.Tags = make(dataset.Tags) + } + series.Header.Tags[promql.MetricNameLabel] = name + } + series.Header.Name = name + series.Header.CalculateHash(true) + series.Header.CalculateSize() + } + } +} + +// Greptime's servers/src/http/prometheus.rs collects metric names through +// aggregates/functions, but clears them for arithmetic and grouping that +// excludes __name__. Regex discovery is expanded upstream and is not static. +func greptimeMetricName(query string) (string, bool) { + expr, err := metricParser.ParseExpr(query) + if err != nil { + return "", false + } + static := true + parser.Inspect(expr, func(node parser.Node, _ []parser.Node) error { + selector, ok := node.(*parser.VectorSelector) + if !ok || selector.Name != "" { + return nil + } + literal := false + for _, matcher := range selector.LabelMatchers { + if matcher.Name == promql.MetricNameLabel && matcher.Type == labels.MatchEqual { + literal = true + } + } + static = static && literal + return nil + }) + if !static { + return "", false + } + names := make(map[string]struct{}) + var collect func(parser.Expr) bool + collect = func(expr parser.Expr) bool { + switch e := expr.(type) { + case *parser.AggregateExpr: + if (e.Without && slices.Contains(e.Grouping, promql.MetricNameLabel)) || + (!e.Without && e.Grouping != nil && !slices.Contains(e.Grouping, promql.MetricNameLabel)) { + clear(names) + return true + } + return collect(e.Expr) + case *parser.UnaryExpr: + clear(names) + case *parser.BinaryExpr: + if e.Op == parser.LAND || e.Op == parser.LOR || e.Op == parser.LUNLESS { + return collect(e.LHS) + } + clear(names) + case *parser.ParenExpr: + return collect(e.Expr) + case *parser.SubqueryExpr: + return collect(e.Expr) + case *parser.MatrixSelector: + return collect(e.VectorSelector) + case *parser.VectorSelector: + if e.Name != "" { + names[e.Name] = struct{}{} + return true + } + for _, m := range e.LabelMatchers { + if m.Name != promql.MetricNameLabel { + continue + } + if m.Type != labels.MatchEqual { + return false + } + names[m.Value] = struct{}{} + return true + } + case *parser.Call: + for _, arg := range e.Args { + if !collect(arg) { + return false + } + } + } + return true + } + if !collect(expr) { + return "", false + } + if len(names) == 1 { + for name := range names { + return name, true + } + } + return "", true +} diff --git a/pkg/backends/influxdb/native_listener.go b/pkg/backends/influxdb/native_listener.go index 5f4864a7a..b4247e9a7 100644 --- a/pkg/backends/influxdb/native_listener.go +++ b/pkg/backends/influxdb/native_listener.go @@ -55,7 +55,7 @@ func NativeListenerAdapter() native.Adapter { return nativeListenerAdapter{} } // SupportsHTTP is true because InfluxDB backends serve their primary HTTP // interface through ordinary HTTP listeners; Flight SQL is an additional // native endpoint. -func (nativeListenerAdapter) SupportsHTTP() bool { return true } +func (nativeListenerAdapter) SupportsHTTP(string) bool { return true } func (nativeListenerAdapter) Protocol() string { return listenerconfig.ProtocolFlightSQL } diff --git a/pkg/backends/influxdb/native_listener_test.go b/pkg/backends/influxdb/native_listener_test.go index bb2a81f07..73e19edb4 100644 --- a/pkg/backends/influxdb/native_listener_test.go +++ b/pkg/backends/influxdb/native_listener_test.go @@ -58,7 +58,7 @@ func TestFlightNativeListenerAdapterContract(t *testing.T) { if adapter.Protocol() != listenerconfig.ProtocolFlightSQL { t.Fatalf("Protocol() = %q", adapter.Protocol()) } - if !adapter.SupportsHTTP() { + if !adapter.SupportsHTTP(providers.InfluxDB) { t.Fatal("SupportsHTTP() = false") } if adapter.Configured(listenerconfig.New("flight")) { diff --git a/pkg/backends/mysql/cache.go b/pkg/backends/mysql/cache.go index 9b2c5e7c8..e70636899 100644 --- a/pkg/backends/mysql/cache.go +++ b/pkg/backends/mysql/cache.go @@ -287,7 +287,7 @@ func (h *protocolHandler) finalizeDeltaResult(merged *sqltypes.Result, if err != nil { return nil, nil, nil, err } - response, err := cropSortedResult(merged, timeIndex, plan.OutputUnit, requested) + response, err := h.cropSortedResult(merged, timeIndex, plan.OutputUnit, requested) if err != nil { return nil, nil, nil, err } @@ -360,7 +360,7 @@ func (h *protocolHandler) collectStreamedResult(session *upstreamSession, return nil, fetchErr } if row == nil { - statusFlags, _, stateErr := originProtocolState(upstream) + statusFlags, _, stateErr := h.originProtocolState(upstream) if stateErr != nil { return nil, stateErr } @@ -393,7 +393,7 @@ func (h *protocolHandler) queryCacheKey(c *vtmysql.Conn, session *upstreamSessio ) string { session.mtx.Lock() database := session.database - timeZone := session.timeZone + timeZone := session.viewLocked().TimeZone collation := session.collation session.mtx.Unlock() var identity strings.Builder @@ -416,7 +416,11 @@ func (h *protocolHandler) queryCacheKey(c *vtmysql.Conn, session *upstreamSessio appendCacheIdentityField(&identity, part) } suffix := checksum.Checksum(identity.String()) - return h.config.BackendName + "." + h.config.CacheKeyPrefix + ".mysql." + engine + "." + suffix + dialect := "" + if h.config.Engine != nil { + dialect = h.dialect() + "." + } + return h.config.BackendName + "." + h.config.CacheKeyPrefix + "." + dialect + "mysql." + engine + "." + suffix } func appendCacheIdentityField(identity *strings.Builder, value string) { @@ -463,7 +467,7 @@ func (h *protocolHandler) mergeResults(parts []*sqltypes.Result, if validateErr := comparator.validateRow(row); validateErr != nil { return nil, validateErr } - epoch, parseErr := resultEpoch(row[timeIndex], plan.OutputUnit) + epoch, parseErr := h.resultEpoch(row[timeIndex], plan.OutputUnit) if parseErr != nil { return nil, parseErr } @@ -541,7 +545,7 @@ func (h *protocolHandler) cropAndSortResult(result *sqltypes.Result, if len(row) <= timeIndex { return nil, errors.New("invalid MySQL delta result row") } - epoch, parseErr := resultEpoch(row[timeIndex], plan.OutputUnit) + epoch, parseErr := h.resultEpoch(row[timeIndex], plan.OutputUnit) if parseErr != nil { return nil, parseErr } @@ -580,14 +584,14 @@ func (h *protocolHandler) cropAndSortResult(result *sqltypes.Result, // cropSortedResult crops a result already ordered by (epoch, group), as // guaranteed by mergeResults, without rebuilding group keys or sorting again. -func cropSortedResult(result *sqltypes.Result, timeIndex int, +func (h *protocolHandler) cropSortedResult(result *sqltypes.Result, timeIndex int, unit timeseries.FieldDataType, extent timeseries.Extent, ) (*sqltypes.Result, error) { - start, err := sortedRowBoundary(result.Rows, timeIndex, unit, extent.Start.UnixNano(), false) + start, err := h.sortedRowBoundary(result.Rows, timeIndex, unit, extent.Start.UnixNano(), false) if err != nil { return nil, err } - end, err := sortedRowBoundary(result.Rows, timeIndex, unit, extent.End.UnixNano(), true) + end, err := h.sortedRowBoundary(result.Rows, timeIndex, unit, extent.End.UnixNano(), true) if err != nil { return nil, err } @@ -599,7 +603,7 @@ func cropSortedResult(result *sqltypes.Result, timeIndex int, return out, nil } -func sortedRowBoundary(rows [][]sqltypes.Value, timeIndex int, +func (h *protocolHandler) sortedRowBoundary(rows [][]sqltypes.Value, timeIndex int, unit timeseries.FieldDataType, target int64, after bool, ) (int, error) { low, high := 0, len(rows) @@ -608,7 +612,7 @@ func sortedRowBoundary(rows [][]sqltypes.Value, timeIndex int, if len(rows[middle]) <= timeIndex { return 0, errors.New("invalid MySQL delta result row") } - epoch, err := resultEpoch(rows[middle][timeIndex], unit) + epoch, err := h.resultEpoch(rows[middle][timeIndex], unit) if err != nil { return 0, err } @@ -719,13 +723,15 @@ type groupColumn struct { // exactly are rejected outright, which costs DPC optimization rather than // correctness because the caller falls back to the object cache. type groupComparator struct { - columns []groupColumn + columns []groupColumn + nullsLast bool } func (h *protocolHandler) newGroupComparator(fields []*querypb.Field, indexes []int, ) (*groupComparator, error) { - c := &groupComparator{columns: make([]groupColumn, len(indexes))} + semantics := h.resultSemantics() + c := &groupComparator{columns: make([]groupColumn, len(indexes)), nullsLast: semantics.NullsLast} for i, index := range indexes { // resultIndexes has already proven every group index addresses a field. field := fields[index] @@ -751,6 +757,10 @@ func (h *protocolHandler) newGroupComparator(fields []*querypb.Field, // deliberately absent: it can be negative, which byte order gets wrong. column.kind = compareBytes case sqltypes.IsText(field.Type): + if semantics.BinaryText { + column.kind = compareBytes + break + } if field.Charset > math.MaxUint16 { return nil, fmt.Errorf("MySQL group column %q uses collation %d, "+ "which Trickster cannot order", field.Name, field.Charset) @@ -807,8 +817,14 @@ func (c *groupComparator) compare(left, right []sqltypes.Value) (int, error) { case l.IsNull() && r.IsNull(): continue case l.IsNull(): + if c.nullsLast { + return 1, nil + } return -1, nil case r.IsNull(): + if c.nullsLast { + return -1, nil + } return 1, nil } order, err := column.compareValues(l, r) @@ -940,7 +956,7 @@ func (h *protocolHandler) applyRetentionSorted(result *sqltypes.Result, if len(row) <= timeIndex { return nil, nil, errors.New("invalid MySQL delta result row") } - epoch, err := resultEpoch(row[timeIndex], plan.OutputUnit) + epoch, err := h.resultEpoch(row[timeIndex], plan.OutputUnit) if err != nil { return nil, nil, err } @@ -1128,11 +1144,11 @@ func (h *protocolHandler) observeAnalysis(statementType sqlparser.StatementType, if counter := h.metricHandles.analysis[key]; counter != nil { counter.Inc() } else { - metrics.SQLQueryAnalysis.WithLabelValues(h.config.BackendName, mysqlDialect, + metrics.SQLQueryAnalysis.WithLabelValues(h.config.BackendName, h.dialect(), analysis.Mode.String(), reason).Inc() } } else { - metrics.SQLQueryAnalysis.WithLabelValues(h.config.BackendName, mysqlDialect, + metrics.SQLQueryAnalysis.WithLabelValues(h.config.BackendName, h.dialect(), analysis.Mode.String(), reason).Inc() } if logger.Level() == level.Debug { @@ -1144,7 +1160,7 @@ func (h *protocolHandler) observeAnalysis(statementType sqlparser.StatementType, } func (h *protocolHandler) observeRewriteFailure(reason string) { - metrics.SQLQueryRewriteFailures.WithLabelValues(h.config.BackendName, mysqlDialect, reason).Inc() + metrics.SQLQueryRewriteFailures.WithLabelValues(h.config.BackendName, h.dialect(), reason).Inc() logger.Error("mysql query extent rewrite failed", logging.Pairs{ keys.BackendName: h.config.BackendName, keys.Reason: reason, }) @@ -1158,7 +1174,7 @@ func (h *protocolHandler) observeCache(mode sqlanalyzer.CacheMode, handles, ok = h.metricHandles.cache[cacheMetricKey{mode: mode, status: status}] } if !ok { - handles = resolveCacheMetricHandles(h.config.BackendName, mode, status) + handles = resolveCacheMetricHandles(h.config.BackendName, h.dialect(), mode, status) } handles.native.Inc() handles.requests.Inc() @@ -1172,7 +1188,7 @@ func (h *protocolHandler) observeCache(mode sqlanalyzer.CacheMode, } } -func newProtocolMetricHandles(backendName string) *protocolMetricHandles { +func newProtocolMetricHandles(backendName, dialect string) *protocolMetricHandles { handles := &protocolMetricHandles{ connectLatency: metrics.MySQLCommandLatency.WithLabelValues(backendName, "connect"), queryLatency: metrics.MySQLCommandLatency.WithLabelValues(backendName, metricPathQuery), @@ -1180,7 +1196,7 @@ func newProtocolMetricHandles(backendName string) *protocolMetricHandles { cache: make(map[cacheMetricKey]cacheMetricHandles, 10), } for _, key := range analysisMetricKeys { - handles.analysis[key] = metrics.SQLQueryAnalysis.WithLabelValues(backendName, mysqlDialect, + handles.analysis[key] = metrics.SQLQueryAnalysis.WithLabelValues(backendName, dialect, key.mode.String(), key.reason) } statuses := map[sqlanalyzer.CacheMode][]cachestatus.LookupStatus{ @@ -1202,13 +1218,13 @@ func newProtocolMetricHandles(backendName string) *protocolMetricHandles { for mode, values := range statuses { for _, status := range values { key := cacheMetricKey{mode: mode, status: status} - handles.cache[key] = resolveCacheMetricHandles(backendName, mode, status) + handles.cache[key] = resolveCacheMetricHandles(backendName, dialect, mode, status) } } return handles } -func resolveCacheMetricHandles(backendName string, mode sqlanalyzer.CacheMode, +func resolveCacheMetricHandles(backendName, dialect string, mode sqlanalyzer.CacheMode, status cachestatus.LookupStatus, ) cacheMetricHandles { httpStatus := metricHTTPStatusOK @@ -1217,13 +1233,13 @@ func resolveCacheMetricHandles(backendName string, mode sqlanalyzer.CacheMode, } statusLabel := status.String() return cacheMetricHandles{ - native: metrics.SQLQueryCache.WithLabelValues(backendName, mysqlDialect, + native: metrics.SQLQueryCache.WithLabelValues(backendName, dialect, mode.String(), statusLabel), - requests: metrics.ProxyRequestStatus.WithLabelValues(backendName, mysqlDialect, + requests: metrics.ProxyRequestStatus.WithLabelValues(backendName, dialect, metricMethodQuery, statusLabel, httpStatus, metricPathQuery), - elements: metrics.ProxyRequestElements.WithLabelValues(backendName, mysqlDialect, + elements: metrics.ProxyRequestElements.WithLabelValues(backendName, dialect, statusLabel, metricPathQuery), - duration: metrics.ProxyRequestDuration.WithLabelValues(backendName, mysqlDialect, + duration: metrics.ProxyRequestDuration.WithLabelValues(backendName, dialect, metricMethodQuery, statusLabel, httpStatus, metricPathQuery), } } diff --git a/pkg/backends/mysql/engine.go b/pkg/backends/mysql/engine.go new file mode 100644 index 000000000..f83d7e442 --- /dev/null +++ b/pkg/backends/mysql/engine.go @@ -0,0 +1,172 @@ +/* + * Copyright 2026 The Trickster Authors + * 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 mysql + +import ( + "errors" + "net" + "net/url" + "strconv" + "time" + + bo "github.com/trickstercache/trickster/v2/pkg/backends/options" + "github.com/trickstercache/trickster/v2/pkg/parsing/sqlanalyzer" + "github.com/trickstercache/trickster/v2/pkg/timeseries" + + vtmysql "vitess.io/vitess/go/mysql" + "vitess.io/vitess/go/sqltypes" + "vitess.io/vitess/go/vt/sqlparser" +) + +// SessionView is the connection-scoped state visible to dialect analysis. +type SessionView struct { + TimeZone string +} + +// Engine supplies a provider's SQL analysis and result-state contract while +// sharing the MySQL listener, authentication, limits and cache transport. +type Engine interface { + Name() string + DefaultPort() string + SupportsHTTP() bool + Analyzer(SessionView) sqlanalyzer.DialectAnalyzer + StreamState(*vtmysql.Conn) (status, warnings uint16, err error) +} + +// SessionInitializer reads effective defaults without modifying the origin session. +type SessionInitializer interface { + InitSession(*vtmysql.Conn) (SessionView, error) +} + +// ResultSemantics supplies the ordering and timestamp rules of a compatible +// engine. Zero values keep MySQL's result rules. +type ResultSemantics struct { + Timestamp func(sqltypes.Value) (int64, error) + BinaryText bool + NullsLast bool +} + +type resultEngine interface{ ResultSemantics() ResultSemantics } + +func (h *protocolHandler) resultSemantics() ResultSemantics { + if e, ok := h.config.Engine.(resultEngine); ok { + return e.ResultSemantics() + } + return ResultSemantics{} +} + +func (h *protocolHandler) resultEpoch(value sqltypes.Value, unit timeseries.FieldDataType) (int64, error) { + if unit == timeseries.DateTimeSQL { + if parser := h.resultSemantics().Timestamp; parser != nil { + return parser(value) + } + } + return resultEpoch(value, unit) +} + +// ProtocolConfigForEngine derives a native endpoint without mutating the HTTP +// origin or the caller's options. A nil engine retains the MySQL defaults. +func ProtocolConfigForEngine(o *bo.Options, engine Engine) (ProtocolConfig, error) { + if engine == nil { + return ProtocolConfigFromOptions(o) + } + copy, err := upstreamOptionsForEngine(o, engine) + if err != nil { + return ProtocolConfig{}, err + } + config, err := ProtocolConfigFromOptions(copy) + if err != nil { + return ProtocolConfig{}, err + } + config.Engine = engine + config.RestartKey = engine.Name() + ":" + config.RestartKey + return config, nil +} + +func upstreamOptionsForEngine(o *bo.Options, engine Engine) (*bo.Options, error) { + if o == nil { + return nil, errors.New("nil MySQL backend options") + } + if engine == nil { + return o, nil + } + copy := o.Clone() + raw := copy.OriginURL + override := copy.MySQL != nil && copy.MySQL.UpstreamURL != "" + if override { + raw = copy.MySQL.UpstreamURL + copy.MySQL.UpstreamURL = "" + } + u, err := url.Parse(raw) + if err != nil { + return nil, errors.New("invalid MySQL upstream URL") + } + if u.Hostname() == "" { + return nil, errors.New("MySQL upstream URL has no host") + } + if !override && engine.SupportsHTTP() && (u.Scheme == "http" || u.Scheme == "https") { + // HTTP credentials must not cross into a native protocol implicitly. + u = &url.URL{Scheme: "mysql", Host: net.JoinHostPort(u.Hostname(), engine.DefaultPort())} + } else if u.Scheme == "mysql" && u.Port() == "" { + u.Host = net.JoinHostPort(u.Hostname(), engine.DefaultPort()) + } + if u.Scheme == "mysql" { + if port, err := strconv.ParseUint(u.Port(), 10, 16); err != nil || port == 0 { + return nil, errors.New("invalid MySQL upstream port") + } + } + copy.OriginURL = u.String() + return copy, nil +} + +func (h *protocolHandler) originProtocolState(upstream *vtmysql.Conn) (uint16, uint16, error) { + if h.config.Engine != nil { + return h.config.Engine.StreamState(upstream) + } + return originProtocolState(upstream) +} + +func (h *protocolHandler) analyzeQuery(query string, parsed parsedQuery, session *upstreamSession, now time.Time) sqlanalyzer.Analysis { + var analyzer sqlanalyzer.DialectAnalyzer = defaultAnalyzer + if h.config.Engine != nil { + session.mtx.Lock() + view := session.viewLocked() + session.mtx.Unlock() + analyzer = h.config.Engine.Analyzer(view) + } + if analyzer == nil { + return sqlanalyzer.Analysis{Reason: sqlanalyzer.ReasonUnsupportedStatement} + } + if a, ok := analyzer.(interface { + AnalyzeParsed(string, sqlparser.Statement, error) sqlanalyzer.Analysis + }); ok { + return a.AnalyzeParsed(query, parsed.statement, parsed.err) + } + return analyzer.Analyze(query, now) +} + +func (session *upstreamSession) viewLocked() SessionView { + zone := session.timeZone + if zone == "" { + zone = session.defaultTimeZone + } + return SessionView{TimeZone: zone} +} + +func (h *protocolHandler) dialect() string { + if h.config.Engine != nil { + return h.config.Engine.Name() + } + return mysqlDialect +} diff --git a/pkg/backends/mysql/engine_test.go b/pkg/backends/mysql/engine_test.go new file mode 100644 index 000000000..a279d6fb4 --- /dev/null +++ b/pkg/backends/mysql/engine_test.go @@ -0,0 +1,103 @@ +/* + * Copyright 2026 The Trickster Authors + * 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 mysql + +import ( + "strings" + "testing" + "time" + + mo "github.com/trickstercache/trickster/v2/pkg/backends/mysql/options" + "github.com/trickstercache/trickster/v2/pkg/parsing/sqlanalyzer" + + vtmysql "vitess.io/vitess/go/mysql" +) + +type testEngine struct{} + +var ( + nativeProtocolConfig = (nativeListenerAdapter{}).nativeProtocolConfig + nativeRouteRuntime = (nativeListenerAdapter{}).nativeRouteRuntime + isNativeRouter = (nativeListenerAdapter{}).isNativeRouter + isNativeUserRouter = (nativeListenerAdapter{}).isNativeUserRouter + isNativeBalancer = (nativeListenerAdapter{}).isNativeBalancer +) + +func (testEngine) Name() string { return "test-engine" } +func (testEngine) DefaultPort() string { return "4002" } +func (testEngine) SupportsHTTP() bool { return true } +func (testEngine) Analyzer(SessionView) sqlanalyzer.DialectAnalyzer { return nil } +func (testEngine) StreamState(*vtmysql.Conn) (uint16, uint16, error) { return 2, 7, nil } + +func TestEngineProtocolConfig(t *testing.T) { + o := validBackendOptions() + o.Provider, o.OriginURL = "test-engine", "http://metrics.example:4000" + o.MySQL = mo.New() + o.MySQL.UpstreamURL = "mysql://origin:password@db.example/database" + c, err := ProtocolConfigForEngine(o, testEngine{}) + if err != nil { + t.Fatal(err) + } + if c.Engine == nil || c.Upstream.Port != 4002 || c.Upstream.Host != "db.example" || c.Upstream.DbName != "database" { + t.Fatalf("wrong native config: %+v", c.Upstream) + } + if o.OriginURL != "http://metrics.example:4000" { + t.Fatal("HTTP origin mutated") + } + before := c.RestartKey + o.MySQL.UpstreamURL += "2" + c, err = ProtocolConfigForEngine(o, testEngine{}) + if err != nil || before == c.RestartKey { + t.Fatalf("upstream change not reflected in restart identity: %v", err) + } + adapter := NewNativeListenerAdapter(testEngine{}) + if !adapter.ServesProvider("mysql") || !adapter.ServesProvider("test-engine") || !adapter.SupportsHTTP("test-engine") || adapter.SupportsHTTP("mysql") { + t.Fatal("provider capabilities were not kept separate") + } +} + +func TestEngineStreamStateAndAnalysis(t *testing.T) { + h := &protocolHandler{config: ProtocolConfig{Engine: testEngine{}}} + status, warnings, err := h.originProtocolState(nil) + if err != nil || status != 2 || warnings != 7 { + t.Fatalf("engine state = %d/%d, %v", status, warnings, err) + } + parsed := parseQuery("SELECT 1") + if a := h.analyzeQuery("SELECT 1", parsed, &upstreamSession{}, time.Now()); a.Mode != sqlanalyzer.CacheModeNone { + t.Fatalf("nil engine analyzer must not use MySQL's analyzer: %+v", a) + } + if (&protocolHandler{}).analyzeQuery("SELECT 1", parsed, &upstreamSession{}, time.Now()).Mode != sqlanalyzer.CacheModeObject { + t.Fatal("default MySQL analyzer changed") + } +} + +func TestEngineUpstreamValidation(t *testing.T) { + for _, tt := range []struct{ name, origin, upstream, want string }{ + {"HTTP credentials stay HTTP", "https://web:secret@db.example:4000", "", "include a username"}, + {"missing host", "http:///metrics", "", "no host"}, + {"port zero", "http://db.example", "mysql://db:secret@db.example:0/public", "invalid MySQL upstream port"}, + {"port overflow", "http://db.example", "mysql://db:secret@db.example:65536/public", "invalid MySQL upstream port"}, + {"wrong override scheme", "http://db.example", "postgres://db:secret@db.example/public", "unsupported MySQL origin scheme"}, + } { + t.Run(tt.name, func(t *testing.T) { + o := validBackendOptions() + o.OriginURL = tt.origin + o.MySQL = mo.New() + o.MySQL.UpstreamURL = tt.upstream + if _, err := ProtocolConfigForEngine(o, testEngine{}); err == nil || !strings.Contains(err.Error(), tt.want) { + t.Fatalf("want %q, got %v", tt.want, err) + } + }) + } +} diff --git a/pkg/backends/mysql/health.go b/pkg/backends/mysql/health.go index a85d2a71a..b327da61b 100644 --- a/pkg/backends/mysql/health.go +++ b/pkg/backends/mysql/health.go @@ -25,6 +25,7 @@ import ( "github.com/trickstercache/trickster/v2/pkg/backends/healthcheck" ho "github.com/trickstercache/trickster/v2/pkg/backends/healthcheck/options" + bo "github.com/trickstercache/trickster/v2/pkg/backends/options" vtmysql "vitess.io/vitess/go/mysql" "vitess.io/vitess/go/mysql/sqlerror" @@ -56,6 +57,20 @@ func (c *Client) HealthCheckProbe() healthcheck.Probe { return newMySQLHealthProbe(params, connectHealthOrigin) } +// HealthCheckProbeForEngine derives origin credentials without requiring the +// downstream authenticator, which may belong to a native user router or pool. +func HealthCheckProbeForEngine(o *bo.Options, engine Engine) (healthcheck.Probe, error) { + options, err := upstreamOptionsForEngine(o, engine) + if err != nil { + return nil, err + } + params, err := upstreamConnParamsFromOptions(options) + if err != nil { + return nil, err + } + return newMySQLHealthProbe(params, connectHealthOrigin), nil +} + func connectHealthOrigin(ctx context.Context, params *vtmysql.ConnParams, ) (healthConnection, error) { diff --git a/pkg/backends/mysql/native_listener.go b/pkg/backends/mysql/native_listener.go index c672b95ec..b9e584e28 100644 --- a/pkg/backends/mysql/native_listener.go +++ b/pkg/backends/mysql/native_listener.go @@ -19,6 +19,7 @@ package mysql import ( "errors" "fmt" + "net/url" "slices" "strconv" "strings" @@ -38,23 +39,53 @@ import ( "github.com/trickstercache/trickster/v2/pkg/proxy/listener/native" ) -type nativeListenerAdapter struct{} +type nativeListenerAdapter struct{ engines map[string]Engine } var _ native.Adapter = nativeListenerAdapter{} // NativeListenerAdapter returns the MySQL implementation of the native // listener extension point. func NativeListenerAdapter() native.Adapter { - return nativeListenerAdapter{} + return &nativeListenerAdapter{} } -func (nativeListenerAdapter) SupportsHTTP() bool { return false } +// NewNativeListenerAdapter serves MySQL and the supplied compatible engines. +func NewNativeListenerAdapter(engines ...Engine) native.Adapter { + a := &nativeListenerAdapter{engines: make(map[string]Engine, len(engines))} + for _, e := range engines { + if e != nil && e.Name() != providers.MySQL { + a.engines[e.Name()] = e + } + } + return a +} + +func (a nativeListenerAdapter) SupportsHTTP(provider string) bool { + e := a.engines[provider] + return e != nil && e.SupportsHTTP() +} func (nativeListenerAdapter) Protocol() string { return listenerconfig.ProtocolMySQL } -func (nativeListenerAdapter) ServesProvider(provider string) bool { return provider == providers.MySQL } +func (a nativeListenerAdapter) ServesProvider(provider string) bool { + return provider == providers.MySQL || a.engines[provider] != nil +} + +func (a nativeListenerAdapter) Providers() []string { + names := []string{providers.MySQL} + for name := range a.engines { + names = append(names, name) + } + slices.Sort(names) + return names +} -func (nativeListenerAdapter) Providers() []string { return []string{providers.MySQL} } +func (a nativeListenerAdapter) configFromOptions(o *bo.Options) (ProtocolConfig, error) { + if o == nil { + return ProtocolConfig{}, errors.New("nil MySQL backend options") + } + return ProtocolConfigForEngine(o, a.engines[strings.ToLower(o.Provider)]) +} func (nativeListenerAdapter) Configured(o *listenerconfig.Options) bool { return o != nil && o.MySQL != nil @@ -70,17 +101,28 @@ func (nativeListenerAdapter) ValidateListener(o *listenerconfig.Options) error { return o.MySQL.Validate() } -func (nativeListenerAdapter) ValidateBackend(o *bo.Options) error { +func (a nativeListenerAdapter) ValidateBackend(o *bo.Options) error { if o != nil && o.MySQL != nil { if err := o.MySQL.Validate(); err != nil { return err } } - _, err := ProtocolConfigFromOptions(o) + if o != nil && a.SupportsHTTP(strings.ToLower(o.Provider)) { + if o.HasHTTPListener { + u, err := url.Parse(o.OriginURL) + if err != nil || u.Hostname() == "" || (u.Scheme != "http" && u.Scheme != "https") { + return errors.New("an HTTP listener requires an http(s) origin_url; use a mysql listener for a MySQL-only backend") + } + } + if !slices.Contains(o.NativeListenerProtocols, listenerconfig.ProtocolMySQL) { + return nil + } + } + _, err := a.configFromOptions(o) return err } -func (nativeListenerAdapter) ValidateUserRouter(c *config.Config, name string, backend *bo.Options) error { +func (a nativeListenerAdapter) ValidateUserRouter(c *config.Config, name string, backend *bo.Options) error { if backend == nil || backend.ALBOptions == nil || backend.ALBOptions.UserRouter == nil { return fmt.Errorf("mysql user router %q has no user-router configuration", name) } @@ -102,7 +144,7 @@ func (nativeListenerAdapter) ValidateUserRouter(c *config.Config, name string, b if terminal == nil { return fmt.Errorf("mysql user router %q references missing backend %q", name, target) } - if terminal.Provider != providers.MySQL { + if !a.ServesProvider(strings.ToLower(terminal.Provider)) { return fmt.Errorf("mysql user router %q target %q must be a direct mysql backend", name, target) } return nil @@ -126,8 +168,8 @@ func (nativeListenerAdapter) ValidateUserRouter(c *config.Config, name string, b return nil } -func (nativeListenerAdapter) ValidateBalancer(c *config.Config, name string, backend *bo.Options) error { - if !isNativeBalancer(c, backend) { +func (a nativeListenerAdapter) ValidateBalancer(c *config.Config, name string, backend *bo.Options) error { + if !a.isNativeBalancer(c, backend) { return fmt.Errorf("mysql load balancer %q requires a mechanism that balances sessions "+ "over a pool of direct mysql backends", name) } @@ -137,8 +179,8 @@ func (nativeListenerAdapter) ValidateBalancer(c *config.Config, name string, bac return nil } -func (nativeListenerAdapter) Describe(c *config.Config, listenerName string) (native.Descriptor, error) { - protocolConfig, _, err := nativeProtocolConfig(c, listenerName) +func (a nativeListenerAdapter) Describe(c *config.Config, listenerName string) (native.Descriptor, error) { + protocolConfig, _, err := a.nativeProtocolConfig(c, listenerName) if err != nil { return native.Descriptor{}, err } @@ -148,8 +190,8 @@ func (nativeListenerAdapter) Describe(c *config.Config, listenerName string) (na return native.Descriptor{RestartKey: protocolConfig.RestartKey}, nil } -func (nativeListenerAdapter) Build(request native.BuildRequest) (listener.ProtocolServer, error) { - protocolConfig, routed, err := nativeProtocolConfig(request.Config, request.ListenerName) +func (a nativeListenerAdapter) Build(request native.BuildRequest) (listener.ProtocolServer, error) { + protocolConfig, routed, err := a.nativeProtocolConfig(request.Config, request.ListenerName) if err != nil { return nil, err } @@ -170,21 +212,21 @@ func (nativeListenerAdapter) Build(request native.BuildRequest) (listener.Protoc if !routed { return NewProtocolServer(*protocolConfig) } - resolver, targets := nativeRouteRuntime(request) + resolver, targets := a.nativeRouteRuntime(request) if resolver == nil || len(targets) == 0 { return nil, errors.New("no usable native route targets") } return NewRoutedProtocolServer(*protocolConfig, resolver, targets) } -func (nativeListenerAdapter) RouteResolver(request native.BuildRequest) backends.RouteResolver { - resolver, _ := nativeRouteRuntime(request) +func (a nativeListenerAdapter) RouteResolver(request native.BuildRequest) backends.RouteResolver { + resolver, _ := a.nativeRouteRuntime(request) return resolver } // nativeProtocolConfig returns the configuration of the single backend mapped // to a MySQL listener; common listener validation guarantees uniqueness. -func nativeProtocolConfig(c *config.Config, listenerName string) (*ProtocolConfig, bool, error) { +func (a nativeListenerAdapter) nativeProtocolConfig(c *config.Config, listenerName string) (*ProtocolConfig, bool, error) { if c == nil { return nil, false, nil } @@ -192,7 +234,7 @@ func nativeProtocolConfig(c *config.Config, listenerName string) (*ProtocolConfi if !o.UsesListener(listenerName) { continue } - if isNativeRouter(c, o) { + if a.isNativeRouter(c, o) { users, err := DownstreamCredentialsFromOptions(o) if err != nil { return nil, false, err @@ -204,10 +246,10 @@ func nativeProtocolConfig(c *config.Config, listenerName string) (*ProtocolConfi if listenerOptions := c.Listeners[listenerName]; listenerOptions != nil { protocolConfig.ApplyListenerOptions(listenerOptions.MySQL) } - protocolConfig.RestartKey = backendName + ":" + routedRestartKey(c, o, users) + protocolConfig.RestartKey = backendName + ":" + a.routedRestartKey(c, o, users) return &protocolConfig, true, nil } - protocolConfig, err := ProtocolConfigFromOptions(o) + protocolConfig, err := a.configFromOptions(o) if listenerOptions := c.Listeners[listenerName]; listenerOptions != nil { protocolConfig.ApplyListenerOptions(listenerOptions.MySQL) } @@ -219,17 +261,17 @@ func nativeProtocolConfig(c *config.Config, listenerName string) (*ProtocolConfi // isNativeRouter reports whether o is an ALB that commits each MySQL session to a backend: // by the user it authenticated as, or by a strategy that balances sessions over a pool -func isNativeRouter(c *config.Config, o *bo.Options) bool { - return isNativeUserRouter(o) || isNativeBalancer(c, o) +func (a nativeListenerAdapter) isNativeRouter(c *config.Config, o *bo.Options) bool { + return a.isNativeUserRouter(o) || a.isNativeBalancer(c, o) } -func isNativeUserRouter(o *bo.Options) bool { +func (a nativeListenerAdapter) isNativeUserRouter(o *bo.Options) bool { return o != nil && o.Provider == providers.ALB && o.ALBOptions != nil && o.ALBOptions.MechanismName == names.MechanismUR && o.ALBOptions.UserRouter != nil && - o.ALBOptions.UserRouter.TargetProvider == providers.MySQL + a.ServesProvider(strings.ToLower(o.ALBOptions.UserRouter.TargetProvider)) } -func isNativeBalancer(c *config.Config, o *bo.Options) bool { +func (a nativeListenerAdapter) isNativeBalancer(c *config.Config, o *bo.Options) bool { if c == nil || o == nil || o.Provider != providers.ALB || o.ALBOptions == nil || o.ALBOptions.UserRouter != nil || len(o.ALBOptions.Pool) == 0 || o.ALBOptions.MechanismName == names.MechanismUR || @@ -237,7 +279,7 @@ func isNativeBalancer(c *config.Config, o *bo.Options) bool { return false } for _, m := range o.ALBOptions.Pool { - if member := c.Backends[m.Name]; member == nil || member.Provider != providers.MySQL { + if member := c.Backends[m.Name]; member == nil || !a.ServesProvider(strings.ToLower(member.Provider)) { return false } } @@ -252,10 +294,10 @@ type nativeRouteProvider interface { MySQLRouteConfig() (ProtocolConfig, error) } -func nativeRouteRuntime(request native.BuildRequest) (backends.RouteResolver, map[string]ProtocolConfig) { +func (a nativeListenerAdapter) nativeRouteRuntime(request native.BuildRequest) (backends.RouteResolver, map[string]ProtocolConfig) { routerName := backendForListener(request.Config, request.ListenerName) routerOptions := request.Config.Backends[routerName] - if !isNativeRouter(request.Config, routerOptions) { + if !a.isNativeRouter(request.Config, routerOptions) { return nil, nil } client := request.BackendClients.Get(routerName) @@ -324,7 +366,7 @@ func routeTargetNames(o *bo.Options) []string { return result } -func routedRestartKey(c *config.Config, router *bo.Options, users map[string]string) string { +func (a nativeListenerAdapter) routedRestartKey(c *config.Config, router *bo.Options, users map[string]string) string { var identity strings.Builder appendRestartIdentityField(&identity, userRouterRestartIdentity(router.ALBOptions.UserRouter)) appendRestartIdentityField(&identity, strconv.FormatBool(router.RequireTLS)) @@ -332,7 +374,7 @@ func routedRestartKey(c *config.Config, router *bo.Options, users map[string]str appendRestartIdentityField(&identity, credentialRestartIdentity(users)) for _, name := range routeTargetNames(router) { if target := c.Backends[name]; target != nil { - if protocolConfig, err := ProtocolConfigFromOptions(target); err == nil { + if protocolConfig, err := a.configFromOptions(target); err == nil { appendRestartIdentityField(&identity, name) appendRestartIdentityField(&identity, protocolConfig.RestartKey) } diff --git a/pkg/backends/mysql/options/options.go b/pkg/backends/mysql/options/options.go index 62c76b162..cf70d010e 100644 --- a/pkg/backends/mysql/options/options.go +++ b/pkg/backends/mysql/options/options.go @@ -48,8 +48,10 @@ const ( // max_concurrent_conns, and max_object_size_bytes options supply the remaining // origin and cache limits. type Options struct { - MaxResultRows int `yaml:"max_result_rows,omitempty"` - MaxResultSizeBytes int `yaml:"max_result_size_bytes,omitempty"` + // UpstreamURL separates the native endpoint from a provider's HTTP origin. + UpstreamURL string `yaml:"upstream_url,omitempty"` + MaxResultRows int `yaml:"max_result_rows,omitempty"` + MaxResultSizeBytes int `yaml:"max_result_size_bytes,omitempty"` } // New returns the default MySQL backend options. diff --git a/pkg/backends/mysql/protocol.go b/pkg/backends/mysql/protocol.go index 46453c009..89987156f 100644 --- a/pkg/backends/mysql/protocol.go +++ b/pkg/backends/mysql/protocol.go @@ -75,6 +75,7 @@ var passwordHashPrefixes = [...]string{ // authenticated upstream session and can grow to include certificate policy // without coupling it to the generic listener package. type ProtocolConfig struct { + Engine Engine BackendName string RestartKey string Upstream vtmysql.ConnParams @@ -163,7 +164,11 @@ func upstreamConnParamsFromOptions(o *bo.Options) (vtmysql.ConnParams, error) { if o == nil { return vtmysql.ConnParams{}, errors.New("nil MySQL backend options") } - u, err := url.Parse(o.OriginURL) + rawURL := o.OriginURL + if o.MySQL != nil && o.MySQL.UpstreamURL != "" { + rawURL = o.MySQL.UpstreamURL + } + u, err := url.Parse(rawURL) if err != nil { return vtmysql.ConnParams{}, fmt.Errorf("parse MySQL origin URL: %w", err) } @@ -431,11 +436,12 @@ func NewRoutedProtocolServer(config ProtocolConfig, resolver backends.RouteResol } func newProtocolHandler(config ProtocolConfig, env *vtenv.Environment) *protocolHandler { - return &protocolHandler{ + h := &protocolHandler{ config: config, env: env, sessions: make(map[*vtmysql.Conn]*upstreamSession), - controls: make(map[uint32]*phaseConn), - metricHandles: newProtocolMetricHandles(config.BackendName), + controls: make(map[uint32]*phaseConn), } + h.metricHandles = newProtocolMetricHandles(config.BackendName, h.dialect()) + return h } // Serve runs the protocol accept loop on l. @@ -589,6 +595,7 @@ type upstreamSession struct { warnings uint16 database string timeZone string + defaultTimeZone string collation collations.ID // effective upstream collation inTx bool cacheUnsafe bool @@ -883,7 +890,7 @@ type protocolHandler struct { func (h *protocolHandler) deltaEngine() *nativedelta.Engine[*sqltypes.Result] { h.deltaOnce.Do(func() { h.delta = nativedelta.New(nativedelta.Config{ - Protocol: mysqlDialect, + Protocol: h.dialect(), BackendName: h.config.BackendName, CacheClient: h.cacheClient, CacheTTL: h.config.CacheTTL, @@ -1054,6 +1061,18 @@ func (h *protocolHandler) connectSession(session *upstreamSession) error { return sqlerror.NewSQLError(sqlerror.CRServerGone, sqlerror.SSNetError, "Trickster could not restore the MySQL origin session") } + if initializer, ok := h.config.Engine.(SessionInitializer); ok { + if h.config.ConnectTimeout > 0 { + _ = conn.GetRawConn().SetDeadline(time.Now().Add(h.config.ConnectTimeout)) + } + view, initErr := initializer.InitSession(conn) + _ = conn.GetRawConn().SetDeadline(time.Time{}) + if initErr != nil { + conn.Close() + return sqlerror.NewSQLError(sqlerror.CRServerGone, sqlerror.SSNetError, "Trickster could not read the origin session defaults") + } + session.defaultTimeZone = view.TimeZone + } if h.closed.Load() { conn.Close() return sqlerror.NewSQLError(sqlerror.CRServerGone, sqlerror.SSUnknownSQLState, @@ -1166,7 +1185,12 @@ func (h *protocolHandler) ComQuery(c *vtmysql.Conn, query string, if parsed.statementType != vtparser.StmtSelect { return h.proxyQuery(session, query, parsed, callback) } - analysis := defaultAnalyzer.AnalyzeParsed(query, parsed.statement, parsed.err) + if _, ok := h.config.Engine.(SessionInitializer); ok { + if err := h.connectSession(session); err != nil { + return err + } + } + analysis := h.analyzeQuery(query, parsed, session, time.Now()) h.observeAnalysis(parsed.statementType, analysis) if h.cacheEligible(session) && analysis.Mode != sqlanalyzer.CacheModeNone { cacheStarted := time.Now() @@ -1518,7 +1542,7 @@ func (h *protocolHandler) streamResultSet(session *upstreamSession, upstream *vt return err } if len(fields) == 0 { - statusFlags, warnings, stateErr := originProtocolState(upstream) + statusFlags, warnings, stateErr := h.originProtocolState(upstream) if stateErr != nil { return stateErr } @@ -1541,7 +1565,7 @@ func (h *protocolHandler) streamResultSet(session *upstreamSession, upstream *vt return fetchErr } if row == nil { - statusFlags, warnings, stateErr := originProtocolState(upstream) + statusFlags, warnings, stateErr := h.originProtocolState(upstream) if stateErr != nil { if emitted && session.downstream != nil { session.downstream.MarkForClose() diff --git a/pkg/backends/options/options.go b/pkg/backends/options/options.go index ff9461f2f..f12d1540b 100644 --- a/pkg/backends/options/options.go +++ b/pkg/backends/options/options.go @@ -271,6 +271,10 @@ type Options struct { // // Name is the Name of the backend, taken from the Key in the Backends Lookup Map Name string `yaml:"-"` + // HasHTTPListener is derived from listener mappings during validation. + HasHTTPListener bool `yaml:"-"` + // NativeListenerProtocols includes direct and ALB-inherited protocol mappings. + NativeListenerProtocols []string `yaml:"-"` // Router is a router.Router containing this backend's Path Routes; it is set during route registration Router router.Router `yaml:"-"` // Scheme is the layer 7 protocol indicator (e.g. 'http'), derived from OriginURL @@ -357,6 +361,7 @@ func (o *Options) Clone() *Options { } out.Hosts = slices.Clone(o.Hosts) out.ListenerNames = slices.Clone(o.ListenerNames) + out.NativeListenerProtocols = slices.Clone(o.NativeListenerProtocols) out.CompressibleTypeList = slices.Clone(o.CompressibleTypeList) if o.CompressibleTypes != nil { out.CompressibleTypes = maps.Clone(o.CompressibleTypes) @@ -980,6 +985,28 @@ func (o *Options) CloneYAMLSafe() *Options { co.OriginURL = parsed.String() } } + if co.Postgres != nil && co.Postgres.UpstreamURL != "" { + parsed, err := url.Parse(co.Postgres.UpstreamURL) + if err != nil { + co.Postgres.UpstreamURL = "*****" + } else if parsed.User != nil { + if _, hasPassword := parsed.User.Password(); hasPassword { + parsed.User = url.UserPassword(parsed.User.Username(), "*****") + co.Postgres.UpstreamURL = parsed.String() + } + } + } + if co.MySQL != nil && co.MySQL.UpstreamURL != "" { + parsed, err := url.Parse(co.MySQL.UpstreamURL) + if err != nil { + co.MySQL.UpstreamURL = "*****" + } else if parsed.User != nil { + if _, hasPassword := parsed.User.Password(); hasPassword { + parsed.User = url.UserPassword(parsed.User.Username(), "*****") + co.MySQL.UpstreamURL = parsed.String() + } + } + } // The runtime default is the backend name, but exporting that implicit // value is noisy and suggests replica_group is relevant to every provider. // Preserve only operator-supplied groupings that differ from the name. diff --git a/pkg/backends/options/options_test.go b/pkg/backends/options/options_test.go index 461c4a243..b3258dbe3 100644 --- a/pkg/backends/options/options_test.go +++ b/pkg/backends/options/options_test.go @@ -1226,3 +1226,46 @@ backends: t.Fatal("clone mutated original postgres options") } } + +func TestPostgresUpstreamURLRedaction(t *testing.T) { + for _, raw := range []string{ + "postgres://user:origin-secret@db.example/public", + "postgres://user:origin-secret%zz@db.example/public", + } { + o := New() + o.Postgres = pgo.New() + o.Postgres.UpstreamURL = raw + if safe := o.ToYAML(); strings.Contains(safe, "origin-secret") { + t.Fatal("YAML exposed pgwire upstream credentials") + } + if o.Postgres.UpstreamURL != raw { + t.Fatal("redaction changed the live configuration") + } + } + for _, raw := range []string{"postgres://db.example/public", "postgres://user@db.example/public"} { + o := New() + o.Postgres = pgo.New() + o.Postgres.UpstreamURL = raw + if o.CloneYAMLSafe().Postgres.UpstreamURL != raw { + t.Fatal("redaction changed a URL without a password") + } + } +} + +func TestMySQLUpstreamURLCloneRedaction(t *testing.T) { + for _, raw := range []string{"mysql://user:origin-secret@db.example/public", "mysql://user:origin-secret%zz@db.example/public"} { + o := New() + o.MySQL = mo.New() + o.MySQL.UpstreamURL = raw + o.NativeListenerProtocols = []string{"mysql", "postgres"} + if strings.Contains(o.ToYAML(), "origin-secret") { + t.Fatal("YAML exposed MySQL upstream credentials") + } + clone := o.Clone() + clone.MySQL.UpstreamURL = "changed" + clone.NativeListenerProtocols[0] = "changed" + if o.MySQL.UpstreamURL != raw || o.NativeListenerProtocols[0] != "mysql" { + t.Fatal("clone mutated live options") + } + } +} diff --git a/pkg/backends/postgres/bucket.go b/pkg/backends/postgres/bucket.go index 6a3b2ed74..3088beac0 100644 --- a/pkg/backends/postgres/bucket.go +++ b/pkg/backends/postgres/bucket.go @@ -21,10 +21,8 @@ import ( "time" "github.com/trickstercache/trickster/v2/pkg/parsing/sqlanalyzer/cockroach" - "github.com/trickstercache/trickster/v2/pkg/timeseries" "github.com/cockroachdb/cockroachdb-parser/pkg/sql/sem/tree" - "github.com/cockroachdb/cockroachdb-parser/pkg/sql/sem/tree/treebin" "github.com/cockroachdb/cockroachdb-parser/pkg/sql/types" ) @@ -33,10 +31,7 @@ const ( fnTimeBucketGapfill = "time_bucket_gapfill" fnDateBin = "date_bin" fnDateTrunc = "date_trunc" - fnFloor = "floor" fnExtract = "extract" - fnDatePart = "date_part" - fieldEpoch = "epoch" // unnamedColumn is the result column name PostgreSQL gives an unaliased expression. unnamedColumn = "?column?" @@ -292,76 +287,9 @@ func dateTrunc(name string, args []tree.Expr) (cockroach.BucketMatch, bool) { } func epochFloor(expr tree.Expr) (cockroach.BucketMatch, bool) { - // matches floor(extract(epoch from column)/N)*N over a timestamp column, and - // floor((column)/N)*N over a column of epoch seconds. Both yield epoch seconds. - none := cockroach.BucketMatch{} - product, ok := unwrapParens(expr).(*tree.BinaryExpr) - if !ok || product.Operator.Symbol != treebin.Mult { - return none, false - } - floored, multiplier := product.Left, product.Right - seconds, ok := positiveInteger(multiplier) - if !ok { - floored, multiplier = multiplier, floored - if seconds, ok = positiveInteger(multiplier); !ok { - return none, false - } - } - floor, ok := unwrapParens(floored).(*tree.FuncExpr) - if !ok || len(floor.Exprs) != 1 || !isPlainCall(floor, fnFloor) { - return none, false - } - quotient, ok := unwrapParens(floor.Exprs[0]).(*tree.BinaryExpr) - if !ok || quotient.Operator.Symbol != treebin.Div { - return none, false - } - // truncating to N and then scaling by anything else is not a bucket - if divisor, ok := positiveInteger(quotient.Right); !ok || divisor != seconds || - seconds > int64((1<<63-1)/time.Second) { - return none, false - } - match := cockroach.BucketMatch{ - Step: time.Duration(seconds) * time.Second, OutputUnit: timeseries.DateTimeUnixSecs, - OutputColumn: unnamedColumn, - } - source := unwrapParens(quotient.Left) - if column, ok := cockroach.ColumnName(source); ok { - match.TimeColumn, match.ColumnUnit = column, timeseries.DateTimeUnixSecs - return match, true - } - epoch, ok := source.(*tree.FuncExpr) - if !ok || len(epoch.Exprs) != 2 || !isPlainCall(epoch, fnExtract) && !isPlainCall(epoch, fnDatePart) { - return none, false - } - if field, ok := epoch.Exprs[0].(*tree.StrVal); !ok || !strings.EqualFold(field.RawString(), fieldEpoch) { - return none, false - } - if match.TimeColumn, ok = cockroach.ColumnName(unwrapParens(epoch.Exprs[1])); !ok { - return none, false - } - return match, true -} - -func isPlainCall(function *tree.FuncExpr, name string) bool { - return function.Filter == nil && function.WindowDef == nil && len(function.OrderBy) == 0 && - function.Type == 0 && strings.EqualFold(function.Func.String(), name) -} - -func positiveInteger(expr tree.Expr) (int64, bool) { - number, ok := unwrapParens(expr).(*tree.NumVal) - if !ok { - return 0, false - } - value, err := number.AsInt64() - return value, err == nil && value > 0 -} - -func unwrapParens(expr tree.Expr) tree.Expr { - for { - paren, ok := expr.(*tree.ParenExpr) - if !ok { - return expr - } - expr = paren.Expr + match, ok := cockroach.EpochFloorMatcher(expr) + if ok { + match.OutputColumn = unnamedColumn } + return match, ok } diff --git a/pkg/backends/postgres/compatibility_test.go b/pkg/backends/postgres/compatibility_test.go index e45068077..d49bd2e72 100644 --- a/pkg/backends/postgres/compatibility_test.go +++ b/pkg/backends/postgres/compatibility_test.go @@ -17,257 +17,26 @@ package postgres import ( - "encoding/json" - "os" - "slices" - "strings" "testing" - "time" "github.com/trickstercache/trickster/v2/pkg/parsing/sqlanalyzer" - "github.com/trickstercache/trickster/v2/pkg/timeseries" + "github.com/trickstercache/trickster/v2/pkg/testutil/sqlcompat" ) -const ( - corpusPath = "testdata/compatibility/v1.json" - corpusZoneUTC = "UTC" - corpusMinimumInterval = "1m" - corpusPolicyNone = "none" - corpusPolicyRange = "range_independent" - corpusPlaceholder = "TRICKSTER_TS" -) - -type compatibilityCorpus struct { - SchemaVersion int `json:"schema_version"` - CorpusVersion string `json:"corpus_version"` - MinimumInterval string `json:"minimum_interval"` - Omissions []string `json:"omissions"` - Cases []compatibilityCase `json:"cases"` -} +const corpusPath = "testdata/compatibility/v1.json" -type compatibilityCase struct { - Name string `json:"name"` - MacroSource []string `json:"macro_source"` - QueryOrigin string `json:"query_origin"` - GrafanaVersion string `json:"grafana_version"` - SessionTimeZone string `json:"session_time_zone"` - ExpandedSQL string `json:"expanded_sql"` - Expected compatibilityExpected `json:"expected"` - Rationale string `json:"rationale"` -} - -type compatibilityExpected struct { - CacheMode string `json:"cache_mode"` - AnalysisReason string `json:"analysis_reason"` - Cadence string `json:"cadence"` - Phase string `json:"phase"` - InputUnit string `json:"input_unit"` - OutputUnit string `json:"output_unit"` - LowerBound string `json:"lower_bound"` - LowerInclusive bool `json:"lower_inclusive"` - UpperBound string `json:"upper_bound"` - UpperInclusive bool `json:"upper_inclusive"` - OpenEnded bool `json:"open_ended"` - OutputColumn string `json:"output_column"` - GroupColumns []string `json:"group_columns"` - CanonicalPolicy string `json:"canonical_policy"` - ExtentRendering bool `json:"extent_rendering"` - CanonicalContains []string `json:"canonical_contains"` - CanonicalExcludes []string `json:"canonical_excludes"` -} - -var ( - corpusModes = map[string]sqlanalyzer.CacheMode{ - "none": sqlanalyzer.CacheModeNone, "object": sqlanalyzer.CacheModeObject, "delta": sqlanalyzer.CacheModeDelta, - } - corpusUnits = map[string]timeseries.FieldDataType{ - "timestamp": timeseries.DateTimeRFC3339Nano, "rfc3339": timeseries.DateTimeRFC3339, - "datetime_sql": timeseries.DateTimeSQL, "unix_seconds": timeseries.DateTimeUnixSecs, - } -) - -func loadCompatibilityCorpus(t testing.TB) compatibilityCorpus { - t.Helper() - data, err := os.ReadFile(corpusPath) - if err != nil { - t.Fatal(err) - } - var corpus compatibilityCorpus - if err := json.Unmarshal(data, &corpus); err != nil { - t.Fatal(err) - } - return corpus -} - -func (c compatibilityCase) analyze(sql string) sqlanalyzer.Analysis { - return testAnalyze(c.SessionTimeZone == corpusZoneUTC, sql) +func corpusAnalyze(zone, sql string) sqlanalyzer.Analysis { + return testAnalyze(zone == "UTC", sql) } func TestCompatibilityCorpus(t *testing.T) { - corpus := loadCompatibilityCorpus(t) - if corpus.SchemaVersion != 1 || corpus.CorpusVersion == "" || corpus.MinimumInterval != corpusMinimumInterval { - t.Fatalf("invalid corpus header: %+v", corpus) - } - if len(corpus.Omissions) == 0 { - t.Fatal("the corpus must record what it deliberately leaves out") - } - seen := make(map[string]struct{}, len(corpus.Cases)) - for _, tc := range corpus.Cases { - t.Run(tc.Name, func(t *testing.T) { - if tc.Name == "" || tc.ExpandedSQL == "" || tc.Rationale == "" || tc.SessionTimeZone == "" || - tc.QueryOrigin == "" || tc.Expected.CanonicalPolicy == "" { - t.Fatalf("case lacks required documentation: %+v", tc) - } - if len(tc.MacroSource) > 0 && tc.GrafanaVersion == "" { - t.Fatal("a Grafana macro case must record the Grafana version it was captured from") - } - if _, ok := seen[tc.Name]; ok { - t.Fatalf("duplicate case %q", tc.Name) - } - seen[tc.Name] = struct{}{} - wantMode, ok := corpusModes[tc.Expected.CacheMode] - if !ok { - t.Fatalf("unknown cache mode %q", tc.Expected.CacheMode) - } - analysis := tc.analyze(tc.ExpandedSQL) - if analysis.Mode != wantMode || string(analysis.Reason) != tc.Expected.AnalysisReason { - t.Fatalf("got %s/%s (%v), want %s/%s", analysis.Mode, analysis.Reason, analysis.Err, - tc.Expected.CacheMode, tc.Expected.AnalysisReason) - } - if wantMode != sqlanalyzer.CacheModeDelta { - if analysis.Plan != nil || tc.Expected.ExtentRendering || tc.Expected.CanonicalPolicy != corpusPolicyNone { - t.Fatalf("a case off the delta path must have no renderable plan: %+v", analysis.Plan) - } - return - } - assertCompatibilityPlan(t, tc, analysis.Plan) - }) - } -} - -func assertCompatibilityPlan(t *testing.T, tc compatibilityCase, plan *sqlanalyzer.QueryPlan) { - t.Helper() - want := tc.Expected - if plan == nil || want.CanonicalPolicy != corpusPolicyRange || !want.ExtentRendering { - t.Fatalf("a delta case needs a plan, a range-independent identity and extent rendering: %+v", want) - } - step, err := time.ParseDuration(want.Cadence) - if err != nil { - t.Fatal(err) - } - phase, err := time.ParseDuration(want.Phase) - if err != nil { - t.Fatal(err) - } - lower, err := time.Parse(time.RFC3339Nano, want.LowerBound) - if err != nil { - t.Fatal(err) - } - inputUnit, inputOK := corpusUnits[want.InputUnit] - outputUnit, outputOK := corpusUnits[want.OutputUnit] - if !inputOK || !outputOK { - t.Fatalf("unknown unit in %q / %q", want.InputUnit, want.OutputUnit) - } - if plan.Step != step || plan.Phase != phase || plan.InputUnit != inputUnit || plan.OutputUnit != outputUnit || - plan.OutputColumn != want.OutputColumn || !slices.Equal(plan.GroupColumns, want.GroupColumns) { - t.Fatalf("plan facts: step %v phase %v in %v out %v column %q groups %v", plan.Step, plan.Phase, - plan.InputUnit, plan.OutputUnit, plan.OutputColumn, plan.GroupColumns) - } - if plan.LowerBound == nil || !plan.LowerBound.Value.Equal(lower) || plan.LowerBound.Inclusive != want.LowerInclusive { - t.Fatalf("lower bound %+v, want %s", plan.LowerBound, want.LowerBound) - } - end := lower.Add(step) - if want.OpenEnded { - if plan.UpperBound != nil { - t.Fatalf("expected an open-ended plan, got upper bound %+v", plan.UpperBound) - } - } else { - upper, err := time.Parse(time.RFC3339Nano, want.UpperBound) - if err != nil { - t.Fatal(err) - } - if plan.UpperBound == nil || !plan.UpperBound.Value.Equal(upper) || plan.UpperBound.Inclusive != want.UpperInclusive { - t.Fatalf("upper bound %+v, want %s", plan.UpperBound, want.UpperBound) - } - // every bound lies on the bucket grid, or partial buckets would be cached as whole ones - if !sqlanalyzer.AlignedToBucket(upper, step, phase) { - t.Fatalf("upper bound %s is off the grid", upper) - } - } - if !sqlanalyzer.AlignedToBucket(lower, step, phase) { - t.Fatalf("lower bound %s is off the grid", lower) - } - for _, fragment := range want.CanonicalContains { - if !strings.Contains(plan.CanonicalSQL, fragment) { - t.Errorf("canonical SQL lacks %q: %s", fragment, plan.CanonicalSQL) - } - } - for _, fragment := range want.CanonicalExcludes { - if strings.Contains(plan.CanonicalSQL, fragment) { - t.Errorf("canonical SQL kept %q: %s", fragment, plan.CanonicalSQL) - } - } - rendered, err := plan.RenderExtent(timeseries.Extent{Start: lower, End: end}) - if err != nil { - t.Fatal(err) - } - if strings.Contains(rendered, corpusPlaceholder) || strings.Contains(rendered, "<$") { - t.Fatalf("rendered SQL kept a placeholder: %s", rendered) - } - // what Trickster sends must mean the same statement: same identity, and exactly the extent asked for - again := tc.analyze(rendered) - if again.Mode != sqlanalyzer.CacheModeDelta || again.Plan.CanonicalSQL != plan.CanonicalSQL { - t.Fatalf("rendering changed the statement:\n%s\n%v / %v", rendered, again.Mode, again.Err) - } - if !again.Plan.LowerBound.Value.Equal(lower) || again.Plan.UpperBound == nil || - !again.Plan.UpperBound.Value.Equal(end.Add(step)) { - t.Fatalf("rendered extent reads back as %v..%v, want %v..%v\n%s", again.Plan.LowerBound.Value, - again.Plan.UpperBound, lower, end.Add(step), rendered) - } + sqlcompat.Run(t, corpusPath, corpusAnalyze) } func TestCompatibilityCorpusCoversGrafanaMacros(t *testing.T) { - corpus := loadCompatibilityCorpus(t) - var sources strings.Builder - for _, tc := range corpus.Cases { - sources.WriteString(strings.Join(tc.MacroSource, "\n")) - sources.WriteByte('\n') - } - for _, macro := range []string{ - "$__time(", "$__timeEpoch(", "$__timeFilter(", "$__timeFrom(", "$__timeTo(", "$__timeGroup(", - "$__timeGroupAlias(", "$__unixEpochFilter(", "$__unixEpochNanoFilter(", "$__unixEpochFrom(", - "$__unixEpochTo(", "$__unixEpochGroup(", "$__unixEpochGroupAlias(", "$__interval", - } { - if !strings.Contains(sources.String(), macro) { - t.Errorf("the corpus does not cover %s", macro) - } - } + sqlcompat.CheckGrafanaMacros(t, corpusPath) } func BenchmarkCompatibilityCorpus(b *testing.B) { - corpus := loadCompatibilityCorpus(b) - for _, tc := range corpus.Cases { - b.Run("Analyze/"+tc.Expected.CacheMode+"/"+tc.Name, func(b *testing.B) { - b.ReportAllocs() - for b.Loop() { - _ = tc.analyze(tc.ExpandedSQL) - } - }) - plan := tc.analyze(tc.ExpandedSQL).Plan - if plan == nil { - continue - } - extent := timeseries.Extent{Start: plan.LowerBound.Value, End: plan.LowerBound.Value.Add(plan.Step)} - b.Run("Render/"+tc.Name, func(b *testing.B) { - b.ReportAllocs() - b.RunParallel(func(pb *testing.PB) { - for pb.Next() { - if _, err := plan.RenderExtent(extent); err != nil { - b.Error(err) - return - } - } - }) - }) - } + sqlcompat.Benchmark(b, corpusPath, corpusAnalyze) } diff --git a/pkg/backends/postgres/guard.go b/pkg/backends/postgres/guard.go index 160b99a14..509bb937a 100644 --- a/pkg/backends/postgres/guard.go +++ b/pkg/backends/postgres/guard.go @@ -16,33 +16,27 @@ package postgres -import ( - "strings" +import "github.com/trickstercache/trickster/v2/pkg/parsing/sqlguard" - "github.com/trickstercache/trickster/v2/pkg/parsing/sqlscan" -) - -type wordClass uint8 +type wordClass = sqlguard.WordClass const ( // wordVolatile is a function whose result changes between identical calls. - wordVolatile wordClass = iota + 1 + wordVolatile = sqlguard.Volatile // wordClock reads the clock; as a time bound it is resolved, anywhere else it is volatile. - wordClock + wordClock = sqlguard.Clock // wordBareClock is a clock function written without parentheses. - wordBareClock + wordBareClock = sqlguard.BareClock // wordUnfaithful is a keyword or type the parser drops or re-spells with another meaning. - wordUnfaithful + wordUnfaithful = sqlguard.Unfaithful // wordGapfill fills empty buckets from the range it is asked for. - wordGapfill + wordGapfill = sqlguard.Gapfill // wordCarry reads values from outside the bucket it is reported in. - wordCarry + wordCarry = sqlguard.Carry // wordSetReturning is a function that yields several rows per call. - wordSetReturning + wordSetReturning = sqlguard.SetReturning ) -const unicodeEscapePrefix = "u&" - var guardedWords = map[string]wordClass{ "random": wordVolatile, "random_normal": wordVolatile, "setseed": wordVolatile, "clock_timestamp": wordVolatile, "timeofday": wordVolatile, @@ -90,58 +84,9 @@ type statementFacts struct { } func scanFacts(sql string) statementFacts { - // finds the guarded words of a statement outside quotes and comments. - // A function name counts only when a call follows it. - var facts statementFacts - scanner := sqlscan.New(sql, sqlscan.Options{}) - pending := wordClass(0) - for { - token, more := scanner.Next() - if !more { - return facts - } - called := token.Kind == sqlscan.Punct && sql[token.Start] == '(' - switch { - case !called: - case pending == wordVolatile: - facts.volatile = true - case pending == wordClock: - facts.clock = true - case pending == wordGapfill: - facts.gapfill = true - case pending == wordCarry: - facts.carries = true - case pending == wordSetReturning: - facts.setReturning = true - } - pending = 0 - switch token.Kind { - case sqlscan.Word: - // the map is keyed in lower case; most words are already - word := sql[token.Start:token.End] - class, ok := guardedWords[word] - if !ok { - class = guardedWords[strings.ToLower(word)] - } - switch class { - case wordBareClock: - facts.clock = true - case wordUnfaithful: - facts.unfaithful = true - default: - pending = class - } - case sqlscan.QuotedIdent: - // a quoted name calls the same function; built-in names are lower case - if class := guardedWords[strings.Trim(sql[token.Start:token.End], `"`)]; class != wordUnfaithful { - pending = class - } - fallthrough - case sqlscan.String: - if token.End-token.Start > len(unicodeEscapePrefix) && - strings.EqualFold(sql[token.Start:token.Start+len(unicodeEscapePrefix)], unicodeEscapePrefix) { - facts.unfaithful = true - } - } + facts := sqlguard.Scan(sql, guardedWords) + return statementFacts{ + volatile: facts.Volatile, clock: facts.Clock, unfaithful: facts.Unfaithful, + gapfill: facts.Gapfill, carries: facts.Carries, setReturning: facts.SetReturning, } } diff --git a/pkg/backends/prometheus/compatible.go b/pkg/backends/prometheus/compatible.go new file mode 100644 index 000000000..cffef6252 --- /dev/null +++ b/pkg/backends/prometheus/compatible.go @@ -0,0 +1,99 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 prometheus + +import ( + "net/http" + "net/url" + "slices" + "strings" + + ho "github.com/trickstercache/trickster/v2/pkg/backends/healthcheck/options" + po "github.com/trickstercache/trickster/v2/pkg/proxy/paths/options" +) + +// Hooks describes the differences of a Prometheus-compatible HTTP API. +// Zero values preserve the Prometheus provider's behavior. +type Hooks struct { + PathPrefix string + CacheKeyParams []string + CacheKeyHeaders []string + // PrepareRequest may normalize provider-specific inputs before cache lookup. + // Returning false relays the original request without caching or rewriting. + PrepareRequest func(*http.Request) bool + HealthCheckConfig func(*url.URL) *ho.Options + // PreserveQueryGrid retains caller timestamps rather than rounding them. + PreserveQueryGrid bool + // AlignQueryGrid rounds range endpoints down to epoch-aligned steps while + // retaining PreserveQueryGrid's millisecond parsing and wire precision. + AlignQueryGrid bool +} + +func pathPrefix(prefix string) string { + if prefix = strings.Trim(prefix, "/"); prefix != "" { + return "/" + prefix + } + return "" +} + +// WithPathPrefix copies paths and prepends the API context root. +func WithPathPrefix(paths po.List, prefix string) po.List { + out := paths.Clone() + for _, p := range out { + p.Path = pathPrefix(prefix) + p.Path + } + return out +} + +// WithCacheKeyParams copies paths and adds query/form cache identity fields. +func WithCacheKeyParams(paths po.List, names ...string) po.List { + out := paths.Clone() + for _, p := range out { + for _, name := range names { + if !slices.Contains(p.CacheKeyParams, name) { + p.CacheKeyParams = append(p.CacheKeyParams, name) + } + } + } + return out +} + +// WithCacheKeyHeaders copies paths and adds case-insensitive header identity. +func WithCacheKeyHeaders(paths po.List, names ...string) po.List { + out := paths.Clone() + for _, p := range out { + for _, name := range names { + if !slices.ContainsFunc(p.CacheKeyHeaders, func(existing string) bool { + return strings.EqualFold(existing, name) + }) { + p.CacheKeyHeaders = append(p.CacheKeyHeaders, http.CanonicalHeaderKey(name)) + } + } + } + return out +} + +// Without copies paths while omitting the specified handler names. +func Without(paths po.List, names ...string) po.List { + out := make(po.List, 0, len(paths)) + for _, p := range paths { + if !slices.Contains(names, p.HandlerName) { + out = append(out, p.Clone()) + } + } + return out +} diff --git a/pkg/backends/prometheus/compatible_test.go b/pkg/backends/prometheus/compatible_test.go new file mode 100644 index 000000000..fa85357ab --- /dev/null +++ b/pkg/backends/prometheus/compatible_test.go @@ -0,0 +1,136 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 prometheus + +import ( + "net/http" + "net/http/httptest" + "net/url" + "slices" + "testing" + "time" + + ho "github.com/trickstercache/trickster/v2/pkg/backends/healthcheck/options" + bo "github.com/trickstercache/trickster/v2/pkg/backends/options" + "github.com/trickstercache/trickster/v2/pkg/proxy/headers" +) + +func TestCompatibleGridAlignment(t *testing.T) { + for _, tc := range []struct { + name, start, end, step string + wantStart, wantEnd int64 + }{ + {"seconds", "1704067207.125", "1704067237.125", "15", 1704067200000, 1704067230000}, + {"milliseconds", "1704067200.125", "1704067201.125", "0.5", 1704067200000, 1704067201000}, + {"before_epoch", "-0.125", "15.125", "7", -7000, 14000}, + {"nondivisor", "20.125", "35.125", "7", 14000, 35000}, + } { + t.Run(tc.name, func(t *testing.T) { + for _, align := range []bool{false, true} { + c, err := NewClientWithHooks("grid", nil, nil, nil, Hooks{PreserveQueryGrid: true, AlignQueryGrid: align}) + if err != nil { + t.Fatal(err) + } + v := url.Values{"query": {"up"}, "start": {tc.start}, "end": {tc.end}, "step": {tc.step}} + trq, _, _, err := c.ParseTimeRangeQuery(httptest.NewRequest("GET", "/api/v1/query_range?"+v.Encode(), nil)) + if err != nil { + t.Fatal(err) + } + start, end := time.UnixMilli(tc.wantStart), time.UnixMilli(tc.wantEnd) + if !align { + start, _ = parseGridTime(tc.start) + end, _ = parseGridTime(tc.end) + } + if !trq.Extent.Start.Equal(start) || !trq.Extent.End.Equal(end) || (align && trq.Phase != 0) { + t.Fatalf("align=%v: extent=%+v phase=%s", align, trq.Extent, trq.Phase) + } + } + }) + } +} + +func TestSupportedPathsIsolation(t *testing.T) { + o := bo.New() + original := SupportedPaths(o) + if o.FastForwardPath != nil { + t.Fatal("path catalogue must not mutate backend options") + } + changed := WithPathPrefix(WithCacheKeyHeaders(WithCacheKeyParams(original, "db", "query"), "x-greptime-db-name"), "/v1/prometheus/") + changed[0].Methods[0] = "PATCH" + changed[1].ResponseHeaders[headers.NameCacheControl] = "private" + if original[0].Path != "/api/v1/query_range" || original[0].Methods[0] != http.MethodGet || slices.Contains(original[0].CacheKeyParams, "db") { + t.Fatal("route transformation mutated its input") + } + if changed[0].Path != "/v1/prometheus/api/v1/query_range" || len(changed[0].CacheKeyParams) != len(original[0].CacheKeyParams)+1 { + t.Fatal("invalid prefix or duplicate cache key parameter") + } + if changed[2].ResponseHeaders[headers.NameCacheControl] == "private" || original[1].ResponseHeaders[headers.NameCacheControl] == "private" { + t.Fatal("route response maps alias each other") + } + filtered := Without(changed, "alerts", "admin", "proxycache", "proxy") + if len(filtered) != 5 || len(changed) != len(original) { + t.Fatalf("wrong supported subset: %d", len(filtered)) + } + filtered[0].CacheKeyHeaders[0] = "changed" + if changed[0].CacheKeyHeaders[0] == "changed" || SupportedPaths(nil)[1].ResponseHeaders[headers.NameCacheControl] == "private" { + t.Fatal("default routes share mutable state") + } +} + +func TestCompatibleHandlerLookup(t *testing.T) { + called := 0 + c, err := NewClientWithHooks("compatible", nil, nil, nil, Hooks{ + PathPrefix: "/prefix", + PrepareRequest: func(*http.Request) bool { called++; return true }, + }) + if err != nil { + t.Fatal(err) + } + lookup := c.HandlerLookup() + lookup["admin"].ServeHTTP(httptest.NewRecorder(), httptest.NewRequest("GET", "/prefix/api/v1/admin", nil)) + if called != 1 || !slices.Contains(c.MergeablePaths(), "/prefix/api/v1/query_range") { + t.Fatal("embedding hooks not applied") + } + delete(lookup, "query") + if c.HandlerLookup()["query"] == nil { + t.Fatal("handler lookup shares mutable map") + } +} + +func TestCompatibleHealthAndOptionIsolation(t *testing.T) { + o := bo.New() + o.OriginURL = "http://example.test/base" + hooks := Hooks{PathPrefix: "/prefix", CacheKeyParams: []string{"db"}, CacheKeyHeaders: []string{"X-Database"}} + c, err := NewClientWithHooks("compatible", o, nil, nil, hooks) + if err != nil { + t.Fatal(err) + } + c.BaseUpstreamURL().Path = "/base" + hooks.CacheKeyParams[0] = "changed" + hooks.CacheKeyHeaders[0] = "Changed" + paths := c.DefaultPathConfigs(o) + if !slices.Contains(paths[0].CacheKeyParams, "db") || !slices.Contains(paths[0].CacheKeyHeaders, "X-Database") { + t.Fatal("constructor retained mutable option slices") + } + if c.DefaultHealthCheckConfig().Path != "/base/prefix/api/v1/query" { + t.Fatalf("unexpected probe: %+v", c.DefaultHealthCheckConfig()) + } + c.hooks.HealthCheckConfig = func(u *url.URL) *ho.Options { h := ho.New(); h.Path = u.Path + "/health"; return h } + if c.DefaultHealthCheckConfig().Path != "/base/health" { + t.Fatal("probe hook not called") + } +} diff --git a/pkg/backends/prometheus/handler_labels.go b/pkg/backends/prometheus/handler_labels.go index e345a054a..742ab2947 100644 --- a/pkg/backends/prometheus/handler_labels.go +++ b/pkg/backends/prometheus/handler_labels.go @@ -42,7 +42,9 @@ func (c *Client) LabelsHandler(w http.ResponseWriter, r *http.Request) { qp, _, _ := params.GetRequestValues(r) // start rounds down, end rounds up — see roundEndTimestampParameterToMinute - roundTimestampsToMinute(qp) + if !c.hooks.PreserveQueryGrid { + roundTimestampsToMinute(qp) + } r.URL = u params.SetRequestValues(r, qp) diff --git a/pkg/backends/prometheus/handler_query.go b/pkg/backends/prometheus/handler_query.go index 7fa1ba6c6..515da55ab 100644 --- a/pkg/backends/prometheus/handler_query.go +++ b/pkg/backends/prometheus/handler_query.go @@ -88,7 +88,7 @@ func (c *Client) QueryHandler(w http.ResponseWriter, r *http.Request) { u := urls.BuildUpstreamURL(r, c.BaseUpstreamURL()) qp, _, _ := params.GetRequestValues(r) // Round time param down to the nearest 15 seconds if it exists - if p := qp.Get(upTime); p != "" { + if p := qp.Get(upTime); p != "" && !c.hooks.PreserveQueryGrid { if i, err := strconv.ParseInt(p, 10, 64); err == nil { qp.Set(upTime, strconv.FormatInt(time.Unix(i, 0).Truncate(c.instantRounder).Unix(), 10)) } diff --git a/pkg/backends/prometheus/handler_series.go b/pkg/backends/prometheus/handler_series.go index 4bb8f8d26..a5de6f46d 100644 --- a/pkg/backends/prometheus/handler_series.go +++ b/pkg/backends/prometheus/handler_series.go @@ -39,7 +39,9 @@ func (c *Client) SeriesHandler(w http.ResponseWriter, r *http.Request) { qp, _, _ := params.GetRequestValues(r) // Round Start and End times down to top of most recent minute for cacheability - roundTimestampsToMinute(qp) + if !c.hooks.PreserveQueryGrid { + roundTimestampsToMinute(qp) + } r.URL = u params.SetRequestValues(r, qp) diff --git a/pkg/backends/prometheus/health.go b/pkg/backends/prometheus/health.go index 0e5160a04..3d15ba689 100644 --- a/pkg/backends/prometheus/health.go +++ b/pkg/backends/prometheus/health.go @@ -22,11 +22,14 @@ import ( // DefaultHealthCheckConfig returns the default HealthCheck Config for this backend provider func (c *Client) DefaultHealthCheckConfig() *ho.Options { + if c.hooks.HealthCheckConfig != nil { + return c.hooks.HealthCheckConfig(c.BaseUpstreamURL()) + } o := ho.New() u := c.BaseUpstreamURL() o.Scheme = u.Scheme o.Host = u.Host - o.Path = u.Path + "/api/v1/query" + o.Path = u.Path + pathPrefix(c.hooks.PathPrefix) + "/api/v1/query" o.Query = "query=up" return o } diff --git a/pkg/backends/prometheus/prometheus.go b/pkg/backends/prometheus/prometheus.go index 683f4b36d..e252abbbf 100644 --- a/pkg/backends/prometheus/prometheus.go +++ b/pkg/backends/prometheus/prometheus.go @@ -22,6 +22,7 @@ import ( "math" "net/http" "net/url" + "slices" "strconv" "strings" "time" @@ -143,6 +144,7 @@ func roundTimestampsToMinute(qp url.Values) { // Client Implements Proxy Client Interface type Client struct { backends.TimeseriesBackend + hooks Hooks instantRounder time.Duration hasTransformations bool injectLabels map[string]string @@ -167,7 +169,16 @@ func NewClient(name string, o *bo.Options, router http.Handler, cache cache.Cache, _ backends.Backends, _ types.Lookup, ) (backends.Backend, error) { - c := &Client{} + return NewClientWithHooks(name, o, router, cache, Hooks{}) +} + +// NewClientWithHooks constructs the Prometheus client for a compatible provider. +func NewClientWithHooks(name string, o *bo.Options, router http.Handler, + cache cache.Cache, hooks Hooks, +) (*Client, error) { + hooks.CacheKeyParams = slices.Clone(hooks.CacheKeyParams) + hooks.CacheKeyHeaders = slices.Clone(hooks.CacheKeyHeaders) + c := &Client{hooks: hooks} b, err := backends.NewTimeseriesBackend(name, o, c.RegisterHandlers, router, cache, modelprom.NewModeler()) c.TimeseriesBackend = b @@ -231,7 +242,11 @@ func (c *Client) ParseTimeRangeQuery(r *http.Request) (*timeseries.TimeRangeQuer if p == "" { return nil, nil, false, errors.MissingURLParam(upStart) } - t, err := parseTime(p) + parse := parseTime + if c.hooks.PreserveQueryGrid { + parse = parseGridTime + } + t, err := parse(p) if err != nil { return nil, nil, false, err } @@ -241,7 +256,7 @@ func (c *Client) ParseTimeRangeQuery(r *http.Request) (*timeseries.TimeRangeQuer if p == "" { return nil, nil, false, errors.MissingURLParam(upEnd) } - t, err = parseTime(p) + t, err = parse(p) if err != nil { return nil, nil, false, err } @@ -252,10 +267,47 @@ func (c *Client) ParseTimeRangeQuery(r *http.Request) (*timeseries.TimeRangeQuer return nil, nil, false, errors.MissingURLParam(upStep) } step, err := parseDuration(p) + if c.hooks.PreserveQueryGrid { + step, err = time.ParseDuration(p + "s") + if err != nil { + step, err = tt.ParseDuration(p) + } + if err == nil && (step <= 0 || step%time.Millisecond != 0) { + err = timeseries.ErrUnknownFormat + } + } if err != nil { return nil, nil, false, err } trq.Step = step + if c.hooks.PreserveQueryGrid { + if trq.Extent.End.Before(trq.Extent.Start) { + return nil, nil, false, timeseries.ErrUnknownFormat + } + if c.hooks.AlignQueryGrid { + for _, at := range []*time.Time{&trq.Extent.Start, &trq.Extent.End} { + remainder := at.UnixNano() % int64(step) + if remainder < 0 { + remainder += int64(step) + } + *at = at.Add(-time.Duration(remainder)) + if !at.Equal(time.Unix(0, at.UnixNano())) { + return nil, nil, false, timeseries.ErrUnknownFormat + } + } + } + trq.Phase = time.Duration(trq.Extent.Start.UnixNano() % step.Nanoseconds()) + if trq.Phase < 0 { + trq.Phase += step + } + trq.CacheKeyElements = map[string]string{"grid_phase_ns": strconv.FormatInt(int64(trq.Phase), 10)} + if trq.Phase != 0 { + // Shared epoch-aligned sharding cannot preserve an offset grid. + if o := c.Configuration(); o != nil && o.DoesShard { + return nil, nil, false, timeseries.ErrUnknownFormat + } + } + } if containsOffsetKeyword(trq.Statement) { trq.IsOffset = true @@ -263,6 +315,9 @@ func (c *Client) ParseTimeRangeQuery(r *http.Request) (*timeseries.TimeRangeQuer } rlo.ExtractFastForwardDisabled(trq.Statement) + if c.hooks.PreserveQueryGrid && (trq.Phase != 0 || step%time.Second != 0) { + rlo.FastForwardDisable = true + } trq.ExtractBackfillTolerance(trq.Statement) if x := strings.Index(trq.Statement, timeseries.BackfillToleranceFlag); x > 1 { @@ -281,6 +336,23 @@ func (c *Client) ParseTimeRangeQuery(r *http.Request) (*timeseries.TimeRangeQuer return trq, rlo, true, nil } +// parseGridTime accepts only the millisecond precision carried by Prometheus +// responses, within the common dataset's nanosecond epoch range. +func parseGridTime(value string) (time.Time, error) { + if v, err := strconv.ParseFloat(value, 64); err == nil { + if math.IsNaN(v) || math.IsInf(v, 0) || math.Abs(v) > float64(math.MaxInt64/int64(time.Second)) || math.Round(v*1000)/1000 != v { + return time.Time{}, timeseries.ErrUnknownFormat + } + } else if v, err := time.Parse(time.RFC3339Nano, value); err == nil && v.Nanosecond()%int(time.Millisecond) != 0 { + return time.Time{}, timeseries.ErrUnknownFormat + } + t, err := parseTime(value) + if err == nil && (t.Year() < 1678 || t.Year() > 2261) { + err = timeseries.ErrUnknownFormat + } + return t, err +} + // parseVectorQuery parses the key parts of an Instantaneous Query from the inbound HTTP Request func parseVectorQuery(r *http.Request, rounder time.Duration) (*timeseries.TimeRangeQuery, error) { trq := ×eries.TimeRangeQuery{Extent: timeseries.Extent{}} diff --git a/pkg/backends/prometheus/routes.go b/pkg/backends/prometheus/routes.go index 14779ee5b..bd53b0d04 100644 --- a/pkg/backends/prometheus/routes.go +++ b/pkg/backends/prometheus/routes.go @@ -31,19 +31,37 @@ import ( ) func (c *Client) RegisterHandlers(handlers.Lookup) { - c.TimeseriesBackend.RegisterHandlers( - handlers.Lookup{ - "health": http.HandlerFunc(c.HealthHandler), - "query_range": http.HandlerFunc(c.QueryRangeHandler), - "query": http.HandlerFunc(c.QueryHandler), - "series": http.HandlerFunc(c.SeriesHandler), - "proxycache": http.HandlerFunc(c.ObjectProxyCacheHandler), - "proxy": http.HandlerFunc(c.ProxyHandler), - "labels": http.HandlerFunc(c.LabelsHandler), - "alerts": http.HandlerFunc(c.AlertsHandler), - "admin": http.HandlerFunc(c.UnsupportedHandler), - }, - ) + c.TimeseriesBackend.RegisterHandlers(c.HandlerLookup()) +} + +// HandlerLookup returns independent handler bindings for an embedding provider. +func (c *Client) HandlerLookup() handlers.Lookup { + lookup := handlers.Lookup{ + "health": http.HandlerFunc(c.HealthHandler), + "query_range": http.HandlerFunc(c.QueryRangeHandler), + "query": http.HandlerFunc(c.QueryHandler), + "series": http.HandlerFunc(c.SeriesHandler), + "proxycache": http.HandlerFunc(c.ObjectProxyCacheHandler), + "proxy": http.HandlerFunc(c.ProxyHandler), + "labels": http.HandlerFunc(c.LabelsHandler), + "alerts": http.HandlerFunc(c.AlertsHandler), + "admin": http.HandlerFunc(c.UnsupportedHandler), + } + if c.hooks.PrepareRequest != nil { + for name, handler := range lookup { + if name == "proxy" || name == "health" { + continue + } + lookup[name] = http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if !c.hooks.PrepareRequest(r) { + c.ProxyHandler(w, r) + return + } + handler.ServeHTTP(w, r) + }) + } + } + return lookup } // MergeablePaths returns the list of Prometheus Paths for which Trickster supports @@ -62,11 +80,27 @@ func MergeablePaths() []string { // MergeablePaths returns the list of Prometheus Paths for which Trickster supports // merging multiple documents into a single response func (c *Client) MergeablePaths() []string { - return MergeablePaths() + paths := MergeablePaths() + for i := range paths { + paths[i] = pathPrefix(c.hooks.PathPrefix) + paths[i] + } + return paths } // DefaultPathConfigs returns the default PathConfigs for the given Provider func (c *Client) DefaultPathConfigs(o *bo.Options) po.List { + paths := WithPathPrefix(SupportedPaths(o), c.hooks.PathPrefix) + paths = WithCacheKeyParams(paths, c.hooks.CacheKeyParams...) + paths = WithCacheKeyHeaders(paths, c.hooks.CacheKeyHeaders...) + if o != nil { + o.FastForwardPath = paths[1].Clone() + } + return paths +} + +// SupportedPaths returns a deep copy of the Prometheus route catalogue. +// It does not mutate the provided backend options. +func SupportedPaths(o *bo.Options) po.List { var rhts map[string]string if o != nil { rhts = map[string]string{ @@ -276,6 +310,5 @@ func (c *Client) DefaultPathConfigs(o *bo.Options) po.List { MatchTypeName: matching.PathMatchNamePrefix, }, } - o.FastForwardPath = paths[1].Clone() - return paths + return paths.Clone() } diff --git a/pkg/backends/prometheus/url.go b/pkg/backends/prometheus/url.go index 89e742946..a36ad5a61 100644 --- a/pkg/backends/prometheus/url.go +++ b/pkg/backends/prometheus/url.go @@ -20,6 +20,7 @@ import ( "net/http" "strconv" "strings" + "time" "github.com/trickstercache/trickster/v2/pkg/proxy/params" "github.com/trickstercache/trickster/v2/pkg/proxy/request" @@ -33,6 +34,10 @@ func (c *Client) SetExtent(r *http.Request, _ *timeseries.TimeRangeQuery, v, _, _ := params.GetRequestValues(r) v.Set(upStart, strconv.FormatInt(extent.Start.Unix(), 10)) v.Set(upEnd, strconv.FormatInt(extent.End.Unix(), 10)) + if c.hooks.PreserveQueryGrid { + v.Set(upStart, extent.Start.UTC().Format(time.RFC3339Nano)) + v.Set(upEnd, extent.End.UTC().Format(time.RFC3339Nano)) + } params.SetRequestValues(r, v) return nil } diff --git a/pkg/backends/providers/providers.go b/pkg/backends/providers/providers.go index 2e3269296..982c4bc45 100644 --- a/pkg/backends/providers/providers.go +++ b/pkg/backends/providers/providers.go @@ -53,6 +53,8 @@ const ( DruidID // Postgres represents the PostgreSQL wire-protocol backend provider PostgresID + // GreptimeDB represents the GreptimeDB backend provider. + GreptimeDBID Backends = "backends" @@ -73,6 +75,7 @@ const ( Graphite = "graphite" Druid = "druid" Postgres = "postgres" + GreptimeDB = "greptimedb" // provider name aliases @@ -119,6 +122,7 @@ var Names = map[string]Provider{ Druid: DruidID, Postgres: PostgresID, TimescaleDB: PostgresID, + GreptimeDB: GreptimeDBID, Proxy: RPID, ReverseProxy: RPID, ReverseProxyShort: RPID, @@ -148,6 +152,7 @@ var supportedTimeSeries = map[string]Provider{ Druid: DruidID, Postgres: PostgresID, TimescaleDB: PostgresID, + GreptimeDB: GreptimeDBID, } // IsSupportedTimeSeriesProvider returns true if the provided time series is supported by Trickster @@ -164,6 +169,7 @@ var supportedHTTPTimeSeries = map[string]Provider{ ClickHouse: ClickHouseID, Graphite: GraphiteID, Druid: DruidID, + GreptimeDB: GreptimeDBID, } // IsSupportedHTTPTimeSeriesProvider returns true if the named provider is a time series @@ -184,15 +190,15 @@ func HTTPTimeSeriesProviderNames() []string { return out } -var supportedTimeSeriesMerge = map[string]Provider{ - Prometheus: PrometheusID, +// IsPrometheusCompatible reports whether the provider exposes Prometheus APIs. +func IsPrometheusCompatible(name string) bool { + return name == Prometheus || name == GreptimeDB } // IsSupportedTimeSeriesMergeProvider returns true if the provided time series is // supported by the Time Series Merge ALB mechanism func IsSupportedTimeSeriesMergeProvider(name string) bool { - _, ok := supportedTimeSeriesMerge[name] - return ok + return IsPrometheusCompatible(name) } func (t Provider) String() string { diff --git a/pkg/backends/providers/registry/registry.go b/pkg/backends/providers/registry/registry.go index 2e117b83d..275cdb980 100644 --- a/pkg/backends/providers/registry/registry.go +++ b/pkg/backends/providers/registry/registry.go @@ -21,6 +21,7 @@ import ( "github.com/trickstercache/trickster/v2/pkg/backends/clickhouse" "github.com/trickstercache/trickster/v2/pkg/backends/druid" "github.com/trickstercache/trickster/v2/pkg/backends/graphite" + "github.com/trickstercache/trickster/v2/pkg/backends/greptimedb" "github.com/trickstercache/trickster/v2/pkg/backends/influxdb" "github.com/trickstercache/trickster/v2/pkg/backends/mysql" "github.com/trickstercache/trickster/v2/pkg/backends/postgres" @@ -41,6 +42,7 @@ func SupportedProviders() types.Lookup { providers.ClickHouse: clickhouse.NewClient, providers.Druid: druid.NewClient, providers.Graphite: graphite.NewClient, + providers.GreptimeDB: greptimedb.NewClient, providers.InfluxDB: influxdb.NewClient, providers.MySQL: mysql.NewClient, providers.Postgres: postgres.NewClient, @@ -60,13 +62,14 @@ func SupportedProviders() types.Lookup { // protocols. Adding another native protocol requires registration here, not a // protocol branch in daemon setup or configuration validation. var nativeListeners = func() native.Registry { - mysqlAdapter := mysql.NativeListenerAdapter() + mysqlAdapter := mysql.NewNativeListenerAdapter(greptimedb.MySQLEngine()) clickhouseAdapter := clickhouse.NativeListenerAdapter() influxdbAdapter := influxdb.NativeListenerAdapter() // Every provider that speaks the PostgreSQL wire protocol shares one // adapter; serving another one is a single engine added to this list. pgwireAdapter := pgwire.NewNativeListenerAdapter(pgwire.NewEngines( postgres.Engine(), + greptimedb.Engine(), )) return native.Registry{ mysqlAdapter.Protocol(): mysqlAdapter, diff --git a/pkg/backends/providers/registry/registry_test.go b/pkg/backends/providers/registry/registry_test.go index 28bab14f5..bf432a6bd 100644 --- a/pkg/backends/providers/registry/registry_test.go +++ b/pkg/backends/providers/registry/registry_test.go @@ -47,12 +47,19 @@ func TestPostgresProvidersShareOneNativeAdapter(t *testing.T) { t.Fatalf("%s serves no providers", protocol) } for _, name := range names { - if !registered.ServesProvider(name) || listeners.GetByProvider(name) != registered { - t.Fatalf("%s lists %q but does not serve it exclusively", protocol, name) + if !registered.ServesProvider(name) || listeners.GetForProvider(protocol, name) != registered { + t.Fatalf("%s lists %q but does not serve it", protocol, name) } } } if listeners.GetByProvider(providers.Prometheus) != nil { t.Fatal("an HTTP-only provider must have no native adapter") } + if supported[providers.GreptimeDB] == nil || listeners.GetForProvider("postgres", providers.GreptimeDB) != adapter || + !adapter.SupportsHTTP(providers.GreptimeDB) || adapter.SupportsHTTP(providers.Postgres) { + t.Fatal("GreptimeDB must share pgwire without changing PostgreSQL's HTTP capability") + } + if listeners.GetByProvider(providers.GreptimeDB) != nil || len(listeners.ForProvider(providers.GreptimeDB)) != 2 { + t.Fatal("GreptimeDB must retain both native adapters without an ambiguous default") + } } diff --git a/pkg/checksum/fnv/fnv.go b/pkg/checksum/fnv/fnv.go index 4c4ed530c..231f41a58 100644 --- a/pkg/checksum/fnv/fnv.go +++ b/pkg/checksum/fnv/fnv.go @@ -45,6 +45,17 @@ func (s *InlineFNV64a) Write(data []byte) (int, error) { return len(data), nil } +// WriteString adds the bytes of str to the running hash without converting it to a []byte. +func (s *InlineFNV64a) WriteString(str string) (int, error) { + hash := uint64(*s) + for i := range len(str) { + hash ^= uint64(str[i]) + hash *= prime64 + } + *s = InlineFNV64a(hash) + return len(str), nil +} + // Sum64 returns the uint64 of the current resulting hash. func (s *InlineFNV64a) Sum64() uint64 { return uint64(*s) diff --git a/pkg/checksum/fnv/fnv_test.go b/pkg/checksum/fnv/fnv_test.go index e8a81537e..c5d01bc3e 100644 --- a/pkg/checksum/fnv/fnv_test.go +++ b/pkg/checksum/fnv/fnv_test.go @@ -30,3 +30,17 @@ func TestInlineFNV64a(t *testing.T) { t.Errorf("unexpected checksum for '%s', wanted %d got %d", input, expected, result) } } + +func TestInlineFNV64aWriteString(t *testing.T) { + h1, h2 := NewInlineFNV64a(), NewInlineFNV64a() + if _, err := h1.Write([]byte("trickster")); err != nil { + t.Fatal(err) + } + n, err := h2.WriteString("trickster") + if err != nil || n != len("trickster") { + t.Fatalf("unexpected WriteString result %d, %v", n, err) + } + if h1.Sum64() != h2.Sum64() { + t.Errorf("WriteString and Write disagree: %d != %d", h2.Sum64(), h1.Sum64()) + } +} diff --git a/pkg/config/testdata/greptimedb.yaml b/pkg/config/testdata/greptimedb.yaml new file mode 100644 index 000000000..9b18a9489 --- /dev/null +++ b/pkg/config/testdata/greptimedb.yaml @@ -0,0 +1,199 @@ +main: + server_name: trickster-test +backends: + default: + timeout: 1m0s + keep_alive_timeout: 2m0s + max_concurrent_conns: 20 + max_idle_conns: 20 + cache_name: default + chunk_read_concurrency_limit: 16 + fetch_concurrency_limit: 16 + chunk_write_concurrency_limit: 16 + healthcheck: {} + timeseries_retention_factor: 1024 + timeseries_eviction_method: oldest + negative_cache_name: default + timeseries_ttl: 6h0m0s + fastforward_ttl: 15s + max_ttl: 25h0m0s + revalidation_factor: 2 + max_object_size_bytes: 524288 + max_capture_bytes: 268435456 + compressible_types: + - text/html + - text/javascript + - text/css + - text/plain + - text/xml + - text/json + - application/json + - application/javascript + - application/xml + tracing_name: default + tls: {} + forwarded_headers: standard + latency_min: 0s + latency_max: 0s + greptime: + provider: greptimedb + listener_names: + - default + - greptime-mysql + - greptime-pg + origin_url: http://greptime.example:4000 + timeout: 1m0s + keep_alive_timeout: 2m0s + max_concurrent_conns: 20 + max_idle_conns: 20 + cache_name: default + cache_key_prefix: greptime.example:4000 + chunk_read_concurrency_limit: 16 + fetch_concurrency_limit: 16 + chunk_write_concurrency_limit: 16 + healthcheck: + interval: 5s + timeout: 3s + timeseries_retention_factor: 1024 + timeseries_eviction_method: oldest + negative_cache_name: default + timeseries_ttl: 6h0m0s + fastforward_ttl: 15s + max_ttl: 25h0m0s + revalidation_factor: 2 + max_object_size_bytes: 524288 + max_capture_bytes: 268435456 + compressible_types: + - text/html + - text/javascript + - text/css + - text/plain + - text/xml + - text/json + - application/json + - application/javascript + - application/xml + tracing_name: default + mysql: + upstream_url: mysql://trickster_ro:%2A%2A%2A%2A%2A@greptime.example:4002/public + max_result_rows: 100000 + max_result_size_bytes: 67108864 + postgres: + upstream_url: postgres://trickster_ro:%2A%2A%2A%2A%2A@greptime.example:4003/public + upstream_tls_mode: disable + max_result_rows: 100000 + max_result_size_bytes: 67108864 + tls: {} + forwarded_headers: standard + authenticator_name: greptime-readers + latency_min: 0s + latency_max: 0s +caches: + default: + provider: memory + index: + reap_interval: 3s + flush_interval: 5s + index_expiry: 8760h0m0s + max_size_bytes: 536870912 + max_size_backoff_bytes: 16777216 + max_size_backoff_objects: 100 + redis: + client_type: standard + protocol: tcp + endpoint: redis:6379 + endpoints: + - redis:6379 + filesystem: + cache_path: /tmp/trickster + bbolt: + filename: trickster.db + bucket: trickster + badger: + directory: /tmp/trickster + value_directory: /tmp/trickster + memory: + max_size_bytes: 536870912 + num_counters: 500000 + timeseries_chunk_factor: 420 + byterange_chunk_size: 4096 +frontend: + listen_port: 8480 + tls_listen_port: 8483 + max_request_body_size_bytes: 10485760 + truncate_request_body_too_large: false + read_header_timeout: 10s +listeners: + default: + port: 8480 + tls_port: 8483 + max_request_body_size_bytes: 10485760 + truncate_request_body_too_large: false + read_header_timeout: 10s + protocol: http + tls_watch_interval: 30s + greptime-mysql: + port: 8491 + max_request_body_size_bytes: 10485760 + truncate_request_body_too_large: false + read_header_timeout: 10s + protocol: mysql + tls_watch_interval: 30s + greptime-pg: + port: 8489 + max_request_body_size_bytes: 10485760 + truncate_request_body_too_large: false + read_header_timeout: 10s + protocol: postgres + tls_watch_interval: 30s + metrics: + port: 8481 + max_request_body_size_bytes: 10485760 + truncate_request_body_too_large: false + read_header_timeout: 10s + protocol: http + tls_watch_interval: 30s + mgmt: + address: 127.0.0.1 + port: 8484 + max_request_body_size_bytes: 10485760 + truncate_request_body_too_large: false + read_header_timeout: 10s + protocol: http + tls_watch_interval: 30s +logging: + log_level: INFO +metrics: + listen_port: 8481 +tracing: + default: + provider: none + service_name: trickster + stdout: {} +negative_caches: + default: {} +mgmt: + listen_address: 127.0.0.1 + listen_port: 8484 + config_handler_path: /trickster/config + config_handler_listener: mgmt + ping_handler_path: /trickster/ping + ready_handler_path: /trickster/ready + health_handler_path: /trickster/health + purge_by_key_path: /trickster/purge/key/ + purge_by_path_path: /trickster/purge/path/ + certificates_handler_path: /trickster/certificates + pprof_listener: both + reload_handler_path: /trickster/config/reload + reload_drain_timeout: 30s + reload_rate_limit: 3s +authenticators: + greptime-readers: + provider: basic + observe_only: false + proxy_preserve: true + users_file: "" + users_file_format: "" + users: + user1: '*****' + config: {} diff --git a/pkg/config/validate/greptimedb_test.go b/pkg/config/validate/greptimedb_test.go new file mode 100644 index 000000000..b9933d664 --- /dev/null +++ b/pkg/config/validate/greptimedb_test.go @@ -0,0 +1,172 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 validate + +import ( + "slices" + "strings" + "testing" + + mo "github.com/trickstercache/trickster/v2/pkg/backends/mysql/options" + bo "github.com/trickstercache/trickster/v2/pkg/backends/options" + "github.com/trickstercache/trickster/v2/pkg/config" + "github.com/trickstercache/trickster/v2/pkg/config/listener" + "github.com/trickstercache/trickster/v2/pkg/config/types" + autho "github.com/trickstercache/trickster/v2/pkg/proxy/authenticator/options" + po "github.com/trickstercache/trickster/v2/pkg/proxy/pgwire/options" +) + +func TestGreptimeDBMultipleNativeProtocols(t *testing.T) { + for _, names := range [][]string{ + {"default"}, {"mysql"}, {"postgres"}, {"mysql", "postgres"}, {"default", "mysql", "postgres"}, + } { + t.Run(strings.Join(names, "+"), func(t *testing.T) { + c := config.NewConfig() + c.Listeners["mysql"] = &listener.Options{Protocol: listener.ProtocolMySQL, ListenPort: 8490} + c.Listeners["postgres"] = &listener.Options{Protocol: listener.ProtocolPostgres, ListenPort: 8491} + o := bo.New() + o.Provider, o.OriginURL, o.ListenerNames = "greptimedb", "http://db.example:4000", names + o.MySQL = mo.New() + o.MySQL.UpstreamURL = "mysql://reader:dev-password@db.example/public" + o.Postgres = po.New() + o.Postgres.UpstreamURL = "postgres://reader:dev-password@db.example/public" + o.AuthenticatorName = "native-clients" + o.AuthOptions = &autho.Options{Users: types.EnvStringMap{"client": "dev-password"}} + c.Backends = bo.Lookup{"greptime": o} + if err := Listeners(c); err != nil { + t.Fatal(err) + } + for _, protocol := range []string{listener.ProtocolMySQL, listener.ProtocolPostgres} { + if slices.Contains(o.NativeListenerProtocols, protocol) != slices.Contains(names, protocol) { + t.Fatalf("protocols %v for listeners %v", o.NativeListenerProtocols, names) + } + } + o.ListenerNames = []string{"default"} + o.MySQL.UpstreamURL = "invalid-unused-native-url" + if err := Listeners(c); err != nil { + t.Fatal(err) + } + if len(o.NativeListenerProtocols) != 0 { + t.Fatal("stale native mapping after revalidation") + } + }) + } +} + +func TestGreptimeDBNativeBalancerProtocols(t *testing.T) { + for _, protocol := range []string{listener.ProtocolMySQL, listener.ProtocolPostgres} { + c := replicaConfig("rr") + c.Listeners["mysql1"].Protocol = protocol + for _, name := range []string{"replica-a", "replica-b"} { + o := c.Backends[name] + o.Provider, o.OriginURL = "greptimedb", "http://db.example:4000" + o.MySQL = mo.New() + o.MySQL.UpstreamURL = "mysql://reader:dev-password@db.example/public" + o.Postgres = po.New() + o.Postgres.UpstreamURL = "postgres://reader:dev-password@db.example/public" + } + if protocol == listener.ProtocolPostgres { + if err := Listeners(c); err == nil || !strings.Contains(err.Error(), "postgres session balancing is not supported") { + t.Fatalf("unsupported postgres session balancing: %v", err) + } + continue + } + if err := Listeners(c); err != nil { + t.Fatal(err) + } + for _, name := range []string{"replica-a", "replica-b"} { + o := c.Backends[name] + if len(o.ListenerNames) != 0 || !slices.Equal(o.NativeListenerProtocols, []string{protocol}) { + t.Fatalf("wrong terminal mappings: %+v", o.NativeListenerProtocols) + } + } + if protocol == listener.ProtocolMySQL { + c.Backends["replica-a"].MySQL.UpstreamURL = "invalid" + } else { + c.Backends["replica-a"].Postgres.UpstreamURL = "invalid" + } + if err := Listeners(c); err == nil { + t.Fatal("native target escaped protocol validation") + } + } +} + +func TestGreptimeDBListenerMapping(t *testing.T) { + for _, names := range [][]string{{"default", "greptime"}, {"greptime"}, {"default"}, nil} { + c := config.NewConfig() + c.Listeners["greptime"] = &listener.Options{Protocol: listener.ProtocolPostgres, ListenPort: 8489} + o := bo.New() + o.Provider, o.OriginURL = "greptimedb", "http://db.example:4000" + o.ListenerNames = names + c.Backends = bo.Lookup{"greptime": o} + if err := Listeners(c); err != nil { + t.Fatalf("listeners %v: %v", names, err) + } + wantHTTP := len(names) == 0 || names[0] == "default" + if o.HasHTTPListener != wantHTTP { + t.Fatalf("listeners %v: HTTP=%t", names, o.HasHTTPListener) + } + // Revalidation must clear the synthesized flag when HTTP is removed. + o.ListenerNames = []string{"greptime"} + if err := Listeners(c); err != nil { + t.Fatal(err) + } + if o.HasHTTPListener { + t.Fatal("stale HTTP listener after revalidation") + } + } +} + +func TestGreptimeDBRejectsInvalidMappings(t *testing.T) { + for _, tt := range []struct { + name, origin string + listeners []string + want string + }{ + {"native URL without listener", "postgres://user:password@db.example/public", nil, "HTTP listener requires"}, + {"native URL on HTTP listener", "postgres://user:password@db.example/public", []string{"default", "greptime"}, "HTTP listener requires"}, + {"bad scheme", "mysql://db.example/public", []string{"greptime"}, "unsupported postgres origin scheme"}, + {"missing host", "http:///v1/sql", []string{"greptime"}, "has no host"}, + {"missing listener", "http://db.example:4000", []string{"missing"}, "undefined listener"}, + } { + t.Run(tt.name, func(t *testing.T) { + o := bo.New() + o.Provider, o.OriginURL, o.ListenerNames = "greptimedb", tt.origin, tt.listeners + c := config.NewConfig() + c.Listeners["greptime"] = &listener.Options{Protocol: listener.ProtocolPostgres, ListenPort: 8489} + c.Backends = bo.Lookup{"greptime": o} + if err := Listeners(c); err == nil || !strings.Contains(err.Error(), tt.want) { + t.Fatalf("want %q, got %v", tt.want, err) + } + }) + } +} + +func TestGreptimeDBDeveloperAuthentication(t *testing.T) { + c, err := config.Load([]string{"-config", "../../../docs/developer/environment/trickster-config/trickster.yaml"}) + if err != nil { + t.Fatal(err) + } + o := c.Backends["greptimedb1"] + if o == nil || o.Postgres == nil || o.Postgres.UpstreamURL == "" { + t.Fatal("developer backend requires a separate pgwire upstream URL") + } + auth := c.Authenticators[o.AuthenticatorName] + if auth == nil || !auth.ProxyPreserve { + t.Fatal("GreptimeDB HTTP origin requires preserved client credentials") + } +} diff --git a/pkg/config/validate/native_protocols.go b/pkg/config/validate/native_protocols.go new file mode 100644 index 000000000..c6603a039 --- /dev/null +++ b/pkg/config/validate/native_protocols.go @@ -0,0 +1,64 @@ +/* + * Copyright 2026 The Trickster Authors + * 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 validate + +import ( + "slices" + "strings" + + "github.com/trickstercache/trickster/v2/pkg/backends/providers" + "github.com/trickstercache/trickster/v2/pkg/config" + "github.com/trickstercache/trickster/v2/pkg/proxy/listener/native" +) + +// nativeBackendProtocols includes protocols inherited from native ALB routes, +// whose terminal backends need no direct listener mapping of their own. +func nativeBackendProtocols(c *config.Config, registry native.Registry) map[string][]string { + result := make(map[string][]string) + for name, backend := range c.Backends { + if backend == nil || backend.IsTemplate { + continue + } + backend.NormalizeListenerNames() + var protocols []string + for _, listenerName := range backend.ListenerNames { + if lo := c.Listeners[listenerName]; lo != nil && registry.Get(strings.ToLower(lo.Protocol)) != nil { + protocols = append(protocols, strings.ToLower(lo.Protocol)) + } + } + result[name] = append(result[name], protocols...) + if backend.Provider != providers.ALB || backend.ALBOptions == nil { + continue + } + if ur := backend.ALBOptions.UserRouter; ur != nil { + if ur.DefaultBackend != "" { + result[ur.DefaultBackend] = append(result[ur.DefaultBackend], protocols...) + } + for _, mapping := range ur.Users { + if mapping != nil && mapping.ToBackend != "" { + result[mapping.ToBackend] = append(result[mapping.ToBackend], protocols...) + } + } + } else { + for _, member := range backend.ALBOptions.Pool { + result[member.Name] = append(result[member.Name], protocols...) + } + } + } + for name, protocols := range result { + slices.Sort(protocols) + result[name] = slices.Compact(protocols) + } + return result +} diff --git a/pkg/config/validate/session_balancer_test.go b/pkg/config/validate/session_balancer_test.go index a809eea7d..ff6798255 100644 --- a/pkg/config/validate/session_balancer_test.go +++ b/pkg/config/validate/session_balancer_test.go @@ -77,7 +77,7 @@ func TestNativeListenerBalancesSessions(t *testing.T) { "foreign member": {"rr", func(c *config.Config) { c.Backends["replica-b"].Provider = providers.ReverseProxyShort c.Backends["replica-b"].OriginURL = "http://example.com" - }, "to be a mysql backend: \"replica-b\" is not"}, + }, "to be a greptimedb or mysql backend: \"replica-b\" is not"}, "missing member": {"rr", func(c *config.Config) { delete(c.Backends, "replica-b") }, "\"replica-b\" is not"}, "discovery": {"rr", func(c *config.Config) { c.Backends["replicas"].ALBOptions.Discovery = &ao.DiscoveryOptions{} diff --git a/pkg/config/validate/validate.go b/pkg/config/validate/validate.go index 6f503f41d..c4b902332 100644 --- a/pkg/config/validate/validate.go +++ b/pkg/config/validate/validate.go @@ -269,6 +269,7 @@ func Listeners(c *config.Config) error { mappedProviders := make(map[string]map[string]string, len(c.Listeners)) nativeListeners := providerregistry.NativeListeners() nativeTargets := nativeUserRouterTargets(c, nativeListeners) + nativeProtocols := nativeBackendProtocols(c, nativeListeners) // the load balancers that serve a stream or native listener; every other one serves requests streamALBs := sets.NewStringSet() for backendName, backend := range c.Backends { @@ -277,18 +278,33 @@ func Listeners(c *config.Config) error { continue } backend.NormalizeListenerNames() - if adapter := nativeListeners.GetByProvider(strings.ToLower(backend.Provider)); adapter != nil { + backend.HasHTTPListener = false + backend.NativeListenerProtocols = nativeProtocols[backendName] + listenerNames := backend.ListenerNames + if len(listenerNames) == 0 && !nativeTargets[backendName] { + listenerNames = []string{listener.DefaultFrontendName} + } + for _, name := range listenerNames { + if lo := c.Listeners[name]; lo != nil && (lo.Protocol == "" || strings.EqualFold(lo.Protocol, listener.ProtocolHTTP)) { + backend.HasHTTPListener = true + } + } + for _, adapter := range nativeListeners.ForProvider(strings.ToLower(backend.Provider)) { if err := adapter.ValidateBackend(backend); err != nil { return fmt.Errorf("%s backend %q: %w", adapter.Protocol(), backendName, err) } - if nativeTargets[backendName] && len(backend.ListenerNames) == 0 { - continue - } + } + if nativeTargets[backendName] && len(backend.ListenerNames) == 0 { + continue } if backend.Provider == providers.ALB && backend.ALBOptions != nil && backend.ALBOptions.UserRouter != nil { targetProvider := strings.ToLower(backend.ALBOptions.UserRouter.TargetProvider) - if adapter := nativeListeners.GetByProvider(targetProvider); adapter != nil { + adapters := nativeListeners.ForProvider(targetProvider) + for _, adapter := range adapters { + if len(adapters) > 1 && !slices.Contains(backend.NativeListenerProtocols, adapter.Protocol()) { + continue + } if err := adapter.ValidateUserRouter(c, backendName, backend); err != nil { return err } @@ -381,9 +397,17 @@ func Listeners(c *config.Config) error { return fmt.Errorf("listener %q with protocol %q cannot map to backend %q with provider %q", name, options.Protocol, backendName, provider) } - if adapter := nativeListeners.GetByProvider(targetProvider); options.Protocol == listener.ProtocolHTTP && adapter != nil && !adapter.SupportsHTTP() { - return fmt.Errorf("backend %q with provider %q requires a listener with protocol %q", - backendName, provider, adapter.Protocol()) + if options.Protocol == listener.ProtocolHTTP { + adapters := nativeListeners.ForProvider(targetProvider) + allowed := len(adapters) == 0 + var protocols []string + for _, adapter := range adapters { + allowed = allowed || adapter.SupportsHTTP(targetProvider) + protocols = append(protocols, adapter.Protocol()) + } + if !allowed { + return fmt.Errorf("backend %q with provider %q requires a listener with protocol %q", backendName, provider, strings.Join(protocols, " or ")) + } } } if options.ListenPort < 0 || options.TLSListenPort < 0 { @@ -665,7 +689,7 @@ func nativeUserRouterTargets(c *config.Config, nativeListeners native.Registry) } continue } - if nativeListeners.GetByProvider(strings.ToLower(backend.ALBOptions.UserRouter.TargetProvider)) == nil { + if len(nativeListeners.ForProvider(strings.ToLower(backend.ALBOptions.UserRouter.TargetProvider))) == 0 { continue } if name := backend.ALBOptions.UserRouter.DefaultBackend; name != "" { diff --git a/pkg/observability/logging/accesslog/middleware.go b/pkg/observability/logging/accesslog/middleware.go index a1bf21487..af4afdc04 100644 --- a/pkg/observability/logging/accesslog/middleware.go +++ b/pkg/observability/logging/accesslog/middleware.go @@ -20,7 +20,6 @@ import ( "context" "encoding/binary" "encoding/hex" - "math/rand/v2" "net" "net/http" "time" @@ -30,6 +29,7 @@ import ( "github.com/trickstercache/trickster/v2/pkg/proxy/headers" "github.com/trickstercache/trickster/v2/pkg/proxy/request" utilmiddleware "github.com/trickstercache/trickster/v2/pkg/util/middleware" + "github.com/trickstercache/trickster/v2/pkg/util/weak/compat" ) // UnmatchedName is the backend and provider name logged for requests that no @@ -160,8 +160,8 @@ func newFields(l *Logger, r *http.Request, w http.ResponseWriter, pathConfig str func newRequestID() string { // a 128-bit identifier correlates log lines and is not a secret, so the fast generator will do var b [16]byte - binary.BigEndian.PutUint64(b[:8], rand.Uint64()) //nolint:gosec // correlation id, not a secret - binary.BigEndian.PutUint64(b[8:], rand.Uint64()) //nolint:gosec // correlation id, not a secret + binary.BigEndian.PutUint64(b[:8], compat.Uint64()) + binary.BigEndian.PutUint64(b[8:], compat.Uint64()) return hex.EncodeToString(b[:]) } diff --git a/pkg/parsing/sqlanalyzer/cockroach/cockroach.go b/pkg/parsing/sqlanalyzer/cockroach/cockroach.go index fbf8ccd51..3c8a07499 100644 --- a/pkg/parsing/sqlanalyzer/cockroach/cockroach.go +++ b/pkg/parsing/sqlanalyzer/cockroach/cockroach.go @@ -214,7 +214,7 @@ func ParseIntervalDuration(s string) (time.Duration, bool) { return 0, false } unit, ok := intervalUnits[fields[i+1]] - if !ok { + if !ok || n > int64((1<<63-1-total)/unit) { return 0, false } total += time.Duration(n) * unit @@ -283,6 +283,10 @@ func DateBinMatcher(name string, args []tree.Expr) (BucketMatch, bool) { if !ok { return BucketMatch{}, false } + return dateBinMatch(step, args) +} + +func dateBinMatch(step time.Duration, args []tree.Expr) (BucketMatch, bool) { column, ok := ColumnName(args[1]) if !ok { return BucketMatch{}, false @@ -435,9 +439,10 @@ func (a *Analyzer) Analyze(statement string, now time.Time) sqlanalyzer.Analysis LowerBound: &sqlanalyzer.Bound{ Value: ranges.lower.value, Inclusive: ranges.lower.inclusive, }, - GroupColumns: groups, - Ordering: ordering, - Renderer: renderer, + GroupColumns: groups, + DropsPartialBuckets: ranges.dropsPartialBuckets, + Ordering: ordering, + Renderer: renderer, } if ranges.upper != nil { plan.UpperBound = &sqlanalyzer.Bound{ @@ -949,12 +954,13 @@ type predicateBound struct { } type rangeAnalysis struct { - lower analyzedBound - upper *analyzedBound - targets []*boundTarget - addSynthetic func(tree.Expr) - timeColumn string - lowerStyle boundStyle + lower analyzedBound + upper *analyzedBound + targets []*boundTarget + addSynthetic func(tree.Expr) + timeColumn string + lowerStyle boundStyle + dropsPartialBuckets bool } func (a *Analyzer) analyzeRanges( @@ -1075,6 +1081,7 @@ func normalizePrimaryBounds( } result.lower.value = sqlanalyzer.CeilBucket(result.lower.value, bucket.step, bucket.phase) rounded = true + result.dropsPartialBuckets = true } } @@ -1112,6 +1119,9 @@ func normalizePrimaryBounds( return ErrUnsafePredicate } result.upper.value = sqlanalyzer.FloorBucket(result.upper.value, bucket.step, bucket.phase) + result.dropsPartialBuckets = true + default: + result.dropsPartialBuckets = true } // col <= X reaches at most the first instant of the bucket holding X, // so that bucket is partial; the floored value is the exclusive @@ -1126,6 +1136,7 @@ func normalizePrimaryBounds( } result.upper.value = sqlanalyzer.FloorBucket(result.upper.value, bucket.step, bucket.phase) rounded = true + result.dropsPartialBuckets = true } result.upper.target.offset = bucket.step } diff --git a/pkg/parsing/sqlanalyzer/cockroach/cockroach_test.go b/pkg/parsing/sqlanalyzer/cockroach/cockroach_test.go index 6795225ff..f7960551e 100644 --- a/pkg/parsing/sqlanalyzer/cockroach/cockroach_test.go +++ b/pkg/parsing/sqlanalyzer/cockroach/cockroach_test.go @@ -520,6 +520,31 @@ func TestRenderExtentIsConcurrent(t *testing.T) { } } +func TestPartialBucketMetadata(t *testing.T) { + for _, test := range []struct { + name, predicate string + drops bool + }{ + {"complete", "time >= 1704067200 AND time < 1704153600", false}, + {"partial_lower", "time >= 1704067207 AND time < 1704153600", true}, + {"partial_upper", "time >= 1704067200 AND time < 1704153607", true}, + {"inclusive_aligned", "time >= 1704067200 AND time <= 1704153600", true}, + {"inclusive_partial", "time >= 1704067200 AND time <= 1704153607", true}, + {"inclusive_complete", "time >= 1704067200 AND time <= 1704153599", false}, + {"open_aligned", "time >= 1704067200", false}, + {"open_partial", "time >= 1704067207", true}, + {"output_discrete", "bucket >= 1704067207 AND bucket <= 1704153607", false}, + } { + t.Run(test.name, func(t *testing.T) { + statement := "SELECT date_bin(INTERVAL '10 seconds', time) AS bucket, avg(v) FROM cpu WHERE " + test.predicate + " GROUP BY 1" + got := newDataFusionAnalyzer().Analyze(statement, time.Time{}) + if got.Mode != sqlanalyzer.CacheModeDelta || got.Plan == nil || got.Plan.DropsPartialBuckets != test.drops { + t.Fatalf("analysis=%+v plan=%+v, drops=%v", got, got.Plan, test.drops) + } + }) + } +} + func rfc3339Literal(v time.Time) string { return "'" + v.UTC().Format(time.RFC3339Nano) + "'" } diff --git a/pkg/parsing/sqlanalyzer/cockroach/compact.go b/pkg/parsing/sqlanalyzer/cockroach/compact.go new file mode 100644 index 000000000..3a79b3643 --- /dev/null +++ b/pkg/parsing/sqlanalyzer/cockroach/compact.go @@ -0,0 +1,75 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 cockroach + +import ( + "time" + + "github.com/cockroachdb/cockroachdb-parser/pkg/sql/sem/tree" +) + +// ParseCompactDuration reads positive integer/unit pairs without allocating. +// Units are case-sensitive and dialect-supplied; signs, fractions, whitespace, +// zero components and overflow are rejected. Only fixed-length units belong in +// the table, which must remain immutable while the parser is in use. +func ParseCompactDuration(s string, units map[string]time.Duration) (time.Duration, bool) { + var total time.Duration + for at := 0; at < len(s); { + start := at + var n int64 + for at < len(s) && s[at] >= '0' && s[at] <= '9' { + digit := int64(s[at] - '0') + if n > ((1<<63-1)-digit)/10 { + return 0, false + } + n = n*10 + digit + at++ + } + if start == at || n == 0 { + return 0, false + } + start = at + for at < len(s) && (s[at] < '0' || s[at] > '9') { + at++ + } + unit, ok := units[s[start:at]] + if !ok || unit <= 0 || n > int64((1<<63-1-total)/unit) { + return 0, false + } + total += time.Duration(n) * unit + } + return total, total > 0 +} + +// CompactDateBinMatcher recognizes date_bin('5m', column [, origin]) using a +// dialect's fixed-length compact units. INTERVAL syntax remains a separate matcher. +func CompactDateBinMatcher(units map[string]time.Duration) BucketMatcher { + return func(name string, args []tree.Expr) (BucketMatch, bool) { + if name != "date_bin" || len(args) < 2 || len(args) > 3 { + return BucketMatch{}, false + } + literal, ok := args[0].(*tree.StrVal) + if !ok { + return BucketMatch{}, false + } + step, ok := ParseCompactDuration(literal.RawString(), units) + if !ok { + return BucketMatch{}, false + } + return dateBinMatch(step, args) + } +} diff --git a/pkg/parsing/sqlanalyzer/cockroach/compact_test.go b/pkg/parsing/sqlanalyzer/cockroach/compact_test.go new file mode 100644 index 000000000..4f937e959 --- /dev/null +++ b/pkg/parsing/sqlanalyzer/cockroach/compact_test.go @@ -0,0 +1,101 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 cockroach + +import ( + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/parsing/sqlanalyzer" +) + +func TestParseCompactDuration(t *testing.T) { + units := map[string]time.Duration{"ns": time.Nanosecond, "ms": time.Millisecond, "s": time.Second, "m": time.Minute, "h": time.Hour, "T": time.Millisecond} + for input, want := range map[string]time.Duration{ + "1ns": time.Nanosecond, "15s": 15 * time.Second, "1h30m": 90 * time.Minute, + "5T": 5 * time.Millisecond, "9223372036854775807ns": 1<<63 - 1, + } { + t.Run(input, func(t *testing.T) { + if got, ok := ParseCompactDuration(input, units); !ok || got != want { + t.Fatalf("got %s/%t, want %s", got, ok, want) + } + }) + } + for _, input := range []string{"", "s", "1", "0s", "+1s", "-1s", "1.5s", "1 s", " 1s", "1s ", "1M", "1month", "1y", "1h0m", "9223372036854775808ns", "9223372036854775807s", "9223372036854775807ns1ns"} { + t.Run(input, func(t *testing.T) { + if _, ok := ParseCompactDuration(input, units); ok { + t.Fatal("accepted invalid or overflowing width") + } + }) + } + if _, ok := ParseCompactDuration("1s", map[string]time.Duration{"s": -1}); ok { + t.Fatal("accepted a negative unit") + } + if allocs := testing.AllocsPerRun(100, func() { _, _ = ParseCompactDuration("1h30m", units) }); allocs != 0 { + t.Fatalf("allocated %v", allocs) + } +} + +func TestCompactDateBinMatcher(t *testing.T) { + a := NewAnalyzer(Options{BucketMatchers: []BucketMatcher{CompactDateBinMatcher(map[string]time.Duration{"m": time.Minute})}, RoundUnalignedTimeBounds: true}) + for _, bucket := range []string{"date_bin('5m', ts)", "date_bin('5m', ts, TIMESTAMP '1969-12-31T23:58:00Z')"} { + sql := "SELECT " + bucket + " AS time, count(*) FROM t WHERE ts >= '2026-01-01T00:00:00Z' AND ts < '2026-01-02T00:00:00Z' GROUP BY 1" + analysis := a.Analyze(sql, time.Now()) + if analysis.Mode != sqlanalyzer.CacheModeDelta || analysis.Plan.Step != 5*time.Minute { + t.Fatalf("%s: %+v", bucket, analysis) + } + } + for _, bucket := range []string{"date_bin('0m', ts)", "date_bin('1M', ts)", "date_bin(width, ts)", "date_bin('5m')", "date_bin('5m', ts, now())", "date_bin('5m', ts + 1)", "date_bin(INTERVAL '5 minutes', ts)"} { + analysis := a.Analyze("SELECT "+bucket+" AS time, count(*) FROM t WHERE ts >= '2026-01-01T00:00:00Z' AND ts < '2026-01-02T00:00:00Z' GROUP BY 1", time.Now()) + if analysis.Mode == sqlanalyzer.CacheModeDelta { + t.Fatalf("accepted %s", bucket) + } + } +} + +func TestIntervalDurationOverflow(t *testing.T) { + for _, input := range []string{"9223372036854775807 seconds", "9223372036854775807 nanoseconds 1 nanosecond", "9223372036854775808 nanoseconds"} { + if _, ok := ParseIntervalDuration(input); ok { + t.Fatalf("accepted overflowing interval %s", input) + } + } + if got, ok := ParseIntervalDuration("9223372036854775807 nanoseconds"); !ok || got != 1<<63-1 { + t.Fatalf("rejected largest fixed interval: %s/%t", got, ok) + } +} + +func FuzzParseCompactDuration(f *testing.F) { + for _, seed := range []string{"1ns", "5m", "1h30m", "500ms", "+1m", "1.5s", "0s", "1month", "9223372036854775807ns", "9223372036854775807ns1ns"} { + f.Add(seed) + } + // The common fixed-length unit subset has an independent standard-library + // oracle. This parser deliberately rejects some spellings that Go accepts. + units := map[string]time.Duration{ + "ns": time.Nanosecond, "us": time.Microsecond, "ms": time.Millisecond, + "s": time.Second, "m": time.Minute, "h": time.Hour, + } + f.Fuzz(func(t *testing.T, input string) { + got, ok := ParseCompactDuration(input, units) + if !ok { + return + } + want, err := time.ParseDuration(input) + if err != nil || got <= 0 || got != want { + t.Fatalf("accepted %q as %v; standard parser: %v, %v", input, got, want, err) + } + }) +} diff --git a/pkg/parsing/sqlanalyzer/cockroach/epoch.go b/pkg/parsing/sqlanalyzer/cockroach/epoch.go new file mode 100644 index 000000000..0557bcb35 --- /dev/null +++ b/pkg/parsing/sqlanalyzer/cockroach/epoch.go @@ -0,0 +1,89 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 cockroach + +import ( + "strings" + "time" + + "github.com/trickstercache/trickster/v2/pkg/timeseries" + + "github.com/cockroachdb/cockroachdb-parser/pkg/sql/sem/tree" + "github.com/cockroachdb/cockroachdb-parser/pkg/sql/sem/tree/treebin" +) + +// EpochFloorMatcher matches floor(extract(epoch FROM column)/N)*N, date_part +// and raw epoch-second variants. The result is also an epoch-second count. +// OutputColumn is left empty because unaliased expression names are dialect-specific. +func EpochFloorMatcher(expr tree.Expr) (BucketMatch, bool) { + none := BucketMatch{} + product, ok := unwrapParens(expr).(*tree.BinaryExpr) + if !ok || product.Operator.Symbol != treebin.Mult { + return none, false + } + floored, multiplier := product.Left, product.Right + seconds, ok := positiveInteger(multiplier) + if !ok { + floored, multiplier = multiplier, floored + if seconds, ok = positiveInteger(multiplier); !ok { + return none, false + } + } + floor, ok := unwrapParens(floored).(*tree.FuncExpr) + if !ok || len(floor.Exprs) != 1 || !plainCall(floor, "floor") { + return none, false + } + quotient, ok := unwrapParens(floor.Exprs[0]).(*tree.BinaryExpr) + if !ok || quotient.Operator.Symbol != treebin.Div { + return none, false + } + if divisor, ok := positiveInteger(quotient.Right); !ok || divisor != seconds || + seconds > int64((1<<63-1)/time.Second) { + return none, false + } + match := BucketMatch{Step: time.Duration(seconds) * time.Second, OutputUnit: timeseries.DateTimeUnixSecs} + source := unwrapParens(quotient.Left) + if column, ok := ColumnName(source); ok { + match.TimeColumn, match.ColumnUnit = column, timeseries.DateTimeUnixSecs + return match, true + } + epoch, ok := source.(*tree.FuncExpr) + if !ok || len(epoch.Exprs) != 2 || !plainCall(epoch, "extract") && !plainCall(epoch, "date_part") { + return none, false + } + if field, ok := epoch.Exprs[0].(*tree.StrVal); !ok || !strings.EqualFold(field.RawString(), "epoch") { + return none, false + } + if match.TimeColumn, ok = ColumnName(unwrapParens(epoch.Exprs[1])); !ok { + return none, false + } + return match, true +} + +func plainCall(function *tree.FuncExpr, name string) bool { + return function.Filter == nil && function.WindowDef == nil && len(function.OrderBy) == 0 && + function.Type == 0 && strings.EqualFold(function.Func.String(), name) +} + +func positiveInteger(expr tree.Expr) (int64, bool) { + number, ok := unwrapParens(expr).(*tree.NumVal) + if !ok { + return 0, false + } + value, err := number.AsInt64() + return value, err == nil && value > 0 +} diff --git a/pkg/parsing/sqlanalyzer/cockroach/epoch_test.go b/pkg/parsing/sqlanalyzer/cockroach/epoch_test.go new file mode 100644 index 000000000..067283e4b --- /dev/null +++ b/pkg/parsing/sqlanalyzer/cockroach/epoch_test.go @@ -0,0 +1,62 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 cockroach + +import ( + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/timeseries" + + "github.com/cockroachdb/cockroachdb-parser/pkg/sql/parser" +) + +func TestEpochFloorMatcher(t *testing.T) { + for _, sql := range []string{ + "floor(extract(epoch FROM ts)/300)*300", "300*floor(date_part('epoch', ts)/300)", + "(floor((ts)/300)*300)", "floor(extract(epoch FROM (ts))/300)*300", + } { + t.Run(sql, func(t *testing.T) { + expr, err := parser.ParseExpr(sql) + if err != nil { + t.Fatal(err) + } + match, ok := EpochFloorMatcher(expr) + if !ok || match.TimeColumn != "ts" || match.Step != 5*time.Minute || match.OutputUnit != timeseries.DateTimeUnixSecs || match.OutputColumn != "" { + t.Fatalf("got %+v/%t", match, ok) + } + }) + } + for _, sql := range []string{ + "ts", "floor(ts/300)+300", "floor(ts/300)*width", "floor(ts/300)*0", "floor(ts/300)*-1", + "floor(ts/300)*0.5", "ceil(ts/300)*300", "floor(ts)*300", "floor(ts+300)*300", + "floor(ts/300)*600", "floor(ts/9223372036854775807)*9223372036854775807", + "floor(extract(day FROM ts)/300)*300", "floor(date_part(1, ts)/300)*300", + "floor(date_part('epoch', ts + 1)/300)*300", "floor(other(ts)/300)*300", "floor(date_part('epoch')/300)*300", + "floor(DISTINCT ts/300)*300", "floor(ts/300) FILTER (WHERE ts > 0)*300", "floor(ts/300) OVER ()*300", + } { + t.Run(sql, func(t *testing.T) { + expr, err := parser.ParseExpr(sql) + if err != nil { + t.Fatal(err) + } + if match, ok := EpochFloorMatcher(expr); ok { + t.Fatalf("accepted %+v", match) + } + }) + } +} diff --git a/pkg/parsing/sqlanalyzer/sqlanalyzer.go b/pkg/parsing/sqlanalyzer/sqlanalyzer.go index b83955367..33452d95f 100644 --- a/pkg/parsing/sqlanalyzer/sqlanalyzer.go +++ b/pkg/parsing/sqlanalyzer/sqlanalyzer.go @@ -154,6 +154,10 @@ type QueryPlan struct { LowerBound *Bound UpperBound *Bound GroupColumns []string + // DropsPartialBuckets reports that range normalization excludes partial + // raw-time buckets. Consumers requiring the original SQL result must use + // object caching or proxying instead of rendering this plan. + DropsPartialBuckets bool // ValueColumns names deterministic numeric result fields consumed as // time-series values. Dialect adapters validate expressions statically and // may validate concrete result types when rows arrive. diff --git a/pkg/parsing/sqlanalyzer/vitess/vitess.go b/pkg/parsing/sqlanalyzer/vitess/vitess.go index c5d62d005..6df66d6ea 100644 --- a/pkg/parsing/sqlanalyzer/vitess/vitess.go +++ b/pkg/parsing/sqlanalyzer/vitess/vitess.go @@ -24,6 +24,7 @@ package vitess import ( "errors" "fmt" + "slices" "strconv" "strings" "time" @@ -61,16 +62,37 @@ var ( // Analyzer converts Vitess's MySQL AST into Trickster's dialect-independent // cache plan. It contains no mutable per-query state and is safe for concurrent use. type Analyzer struct { - parser *sqlparser.Parser + parser *sqlparser.Parser + buckets []BucketMatcher + functions map[string]struct{} +} + +// BucketMatcher recognizes an epoch-aligned, fixed-width bucket over one column. +// A dialect must also establish the origin's semantics for any added functions. +type BucketMatcher func(sqlparser.Expr) (*sqlparser.ColName, time.Duration, timeseries.FieldDataType, bool) + +// Options adds dialect-specific syntax without changing MySQL's defaults. +type Options struct { + BucketMatchers []BucketMatcher + DeterministicFunctions []string } // NewAnalyzer returns an analyzer configured for MySQL 8 syntax. func NewAnalyzer() (*Analyzer, error) { + return NewAnalyzerWithOptions(Options{}) +} + +// NewAnalyzerWithOptions returns a parser with optional compatible-dialect rules. +func NewAnalyzerWithOptions(options Options) (*Analyzer, error) { p, err := sqlparser.New(sqlparser.Options{MySQLServerVersion: "8.0.0"}) if err != nil { return nil, err } - return &Analyzer{parser: p}, nil + functions := make(map[string]struct{}, len(options.DeterministicFunctions)) + for _, name := range options.DeterministicFunctions { + functions[strings.ToLower(name)] = struct{}{} + } + return &Analyzer{parser: p, buckets: slices.Clone(options.BucketMatchers), functions: functions}, nil } var _ sqlanalyzer.DialectAnalyzer = (*Analyzer)(nil) @@ -92,6 +114,7 @@ func (a *Analyzer) Parser() *sqlparser.Parser { } type bucketInfo struct { + expression sqlparser.Expr timeColumn string timeAxis string outputColumn string @@ -157,7 +180,7 @@ func (a *Analyzer) AnalyzeParsed(statement string, stmt sqlparser.Statement, } } if selectStmt.Cache != nil || selectStmt.Lock != sqlparser.NoLock || selectStmt.SQLCalcFoundRows || - selectStmt.Into != nil || isNondeterministic(selectStmt) { + selectStmt.Into != nil || a.isNondeterministic(selectStmt) { return sqlanalyzer.Analysis{ Mode: sqlanalyzer.CacheModeNone, Reason: sqlanalyzer.ReasonNondeterministic, Err: ErrUnsupportedStatement, @@ -170,7 +193,7 @@ func (a *Analyzer) AnalyzeParsed(statement string, stmt sqlparser.Statement, len(selectStmt.Windows) > 0 || containsSubquery(selectStmt) { return sqlanalyzer.ObjectAnalysis(sqlanalyzer.ReasonUnsupportedFormat, ErrUnsupportedResultShape) } - bucket, err := analyzeBucket(selectStmt) + bucket, err := analyzeBucket(selectStmt, a.buckets...) if err != nil { return sqlanalyzer.ObjectAnalysis(sqlanalyzer.ReasonUnsupportedBucket, err) } @@ -232,12 +255,15 @@ func containsSubquery(stmt sqlparser.SQLNode) bool { return found } -func isNondeterministic(stmt sqlparser.SQLNode) bool { +func (a *Analyzer) isNondeterministic(stmt sqlparser.SQLNode) bool { unsafe := false _ = sqlparser.Walk(func(node sqlparser.SQLNode) (bool, error) { switch n := node.(type) { case *sqlparser.FuncExpr: name := strings.ToLower(n.Name.String()) + if _, allowed := a.functions[name]; allowed && n.Qualifier.IsEmpty() { + return true, nil + } switch name { case fromUnixTimeFunction, "coalesce", "floor", "ifnull", "round": // These are the deterministic general functions used by supported @@ -335,7 +361,7 @@ func sqlCommentText(statement string) string { return comments.String() } -func analyzeBucket(stmt *sqlparser.Select) (bucketInfo, error) { +func analyzeBucket(stmt *sqlparser.Select, matchers ...BucketMatcher) (bucketInfo, error) { if stmt.SelectExprs == nil { return bucketInfo{}, ErrUnsupportedBucket } @@ -346,6 +372,17 @@ func analyzeBucket(stmt *sqlparser.Select) (bucketInfo, error) { continue } column, axis, seconds, unit, ok := matchBucketExpr(ae.Expr) + if !ok { + for _, match := range matchers { + col, step, outputUnit, matched := match(ae.Expr) + if !matched || col == nil || step <= 0 { + continue + } + column, axis, ok = columnReference(col) + seconds, unit = step, outputUnit + break + } + } if !ok { continue } @@ -357,6 +394,7 @@ func analyzeBucket(stmt *sqlparser.Select) (bucketInfo, error) { return bucketInfo{}, ErrUnsupportedBucket } candidate := bucketInfo{ + expression: ae.Expr, timeColumn: column, outputColumn: alias, timeAxis: axis, step: seconds, unit: unit, } @@ -756,13 +794,13 @@ func selectOutputs(stmt *sqlparser.Select, bucket bucketInfo) ([]selectOutput, i return nil, -1, ErrUnsupportedResultShape } seenNames[key] = struct{}{} - _, axis, _, _, isBucket := matchBucketExpr(aliased.Expr) + isBucket := aliased.Expr == bucket.expression output := selectOutput{ expr: aliased.Expr, name: name, alias: alias, sourceName: sourceName, sourceAxis: sourceAxis, bucket: isBucket, } if isBucket { - if bucketIndex >= 0 || !strings.EqualFold(axis, bucket.timeAxis) || + if bucketIndex >= 0 || !strings.EqualFold(name, bucket.outputColumn) { return nil, -1, ErrUnsupportedResultShape } diff --git a/pkg/parsing/sqlguard/guard.go b/pkg/parsing/sqlguard/guard.go new file mode 100644 index 000000000..3a21895b0 --- /dev/null +++ b/pkg/parsing/sqlguard/guard.go @@ -0,0 +1,112 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 sqlguard finds dialect-defined function and keyword hazards without +// depending on whether a statement is accepted by an AST parser. +package sqlguard + +import ( + "strings" + + "github.com/trickstercache/trickster/v2/pkg/parsing/sqlscan" +) + +// WordClass describes the effect of a guarded word. +type WordClass uint8 + +const ( + // Volatile marks a call whose result or side effects vary between requests. + Volatile WordClass = iota + 1 + // Clock marks a clock-reading call, which may be resolved in a time bound. + Clock + // BareClock marks a clock read written without parentheses. + BareClock + // Unfaithful marks a keyword or type the formatter cannot preserve. + Unfaithful + // Gapfill marks a call that generates missing buckets from the requested range. + Gapfill + // Carry marks a call that depends on values outside its bucket. + Carry + // SetReturning marks a call that can emit multiple rows. + SetReturning +) + +// Facts summarizes hazards outside strings and comments. +type Facts struct { + Volatile bool + Clock bool + Unfaithful bool + Gapfill bool + Carries bool + SetReturning bool +} + +const unicodeEscapePrefix = "u&" + +// Scan reads words using a dialect's immutable lower-case table. Calls require +// an opening parenthesis; quoted lower-case function names are recognized too. +func Scan(sql string, words map[string]WordClass) Facts { + var facts Facts + scanner := sqlscan.New(sql, sqlscan.Options{}) + pending := WordClass(0) + for { + token, more := scanner.Next() + if !more { + return facts + } + if token.Kind == sqlscan.Punct && sql[token.Start] == '(' { + switch pending { + case Volatile: + facts.Volatile = true + case Clock: + facts.Clock = true + case Gapfill: + facts.Gapfill = true + case Carry: + facts.Carries = true + case SetReturning: + facts.SetReturning = true + } + } + pending = 0 + switch token.Kind { + case sqlscan.Word: + word := sql[token.Start:token.End] + class, ok := words[word] + if !ok { + class = words[strings.ToLower(word)] + } + switch class { + case BareClock: + facts.Clock = true + case Unfaithful: + facts.Unfaithful = true + default: + pending = class + } + case sqlscan.QuotedIdent: + if class := words[strings.Trim(sql[token.Start:token.End], `"`)]; class != Unfaithful { + pending = class + } + fallthrough + case sqlscan.String: + if token.End-token.Start > len(unicodeEscapePrefix) && + strings.EqualFold(sql[token.Start:token.Start+len(unicodeEscapePrefix)], unicodeEscapePrefix) { + facts.Unfaithful = true + } + } + } +} diff --git a/pkg/parsing/sqlguard/guard_test.go b/pkg/parsing/sqlguard/guard_test.go new file mode 100644 index 000000000..7871b5345 --- /dev/null +++ b/pkg/parsing/sqlguard/guard_test.go @@ -0,0 +1,46 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 sqlguard + +import "testing" + +func TestScan(t *testing.T) { + words := map[string]WordClass{ + "random": Volatile, "now": Clock, "current_timestamp": BareClock, + "only": Unfaithful, "gapfill": Gapfill, "locf": Carry, "unnest": SetReturning, + } + for sql, want := range map[string]Facts{ + "SELECT RANDOM /* between */ ()": {Volatile: true}, + `SELECT "random"()`: {Volatile: true}, + "SELECT schema.random()": {Volatile: true}, + "SELECT random, 'random()', $$now()$$ /* current_timestamp */": {}, + `SELECT "RANDOM"(), "only", "current_timestamp"`: {}, + "SELECT now()": {Clock: true}, "SELECT current_timestamp": {Clock: true}, + "SELECT * FROM ONLY t": {Unfaithful: true}, "SELECT U&'a'": {Unfaithful: true}, + `SELECT U&"a"`: {Unfaithful: true}, + "SELECT gapfill(), locf(), unnest()": {Gapfill: true, Carries: true, SetReturning: true}, + } { + t.Run(sql, func(t *testing.T) { + if got := Scan(sql, words); got != want { + t.Fatalf("got %+v, want %+v", got, want) + } + }) + } + if got := Scan("SELECT random()", nil); got != (Facts{}) { + t.Fatal("word table was not dialect-specific") + } +} diff --git a/pkg/proxy/engines/deltaproxycache.go b/pkg/proxy/engines/deltaproxycache.go index 5a13de709..ca0f8172c 100644 --- a/pkg/proxy/engines/deltaproxycache.go +++ b/pkg/proxy/engines/deltaproxycache.go @@ -132,6 +132,33 @@ func fetchFastForward( return ffStatus } +// prepareDPCResponse validates before cache or client writes. When fast-forward +// cannot change the data and rendering is request-independent, keep the bytes +// instead of discarding a complete serialization and repeating it later. +func prepareDPCResponse(rts timeseries.Timeseries, rlo *timeseries.RequestOptions, + modeler *timeseries.Modeler, statusCode int, +) ([]byte, error) { + if !rlo.FallbackToProxyOnError { + return nil, nil + } + if !rlo.FastForwardDisable || rlo.MarshalVariesByRequest { + return nil, modeler.WireMarshalWriter(rts, rlo, statusCode, io.Discard) + } + extents := rts.Extents() + rts.SetExtents(nil) + var buf bytes.Buffer + err := modeler.WireMarshalWriter(rts, rlo, statusCode, &buf) + rts.SetExtents(extents) + if err != nil { + return nil, err + } + body := buf.Bytes() + if body == nil { + body = []byte{} + } + return body, nil +} + // finalizeDPCResponse writes metrics, logs, and the HTTP response for a DPC request. // If wireBody is non-nil, it is written directly (skipping marshal). // Otherwise rts is marshaled to the wire format. @@ -355,6 +382,9 @@ func DeltaProxyCacheRequest(w http.ResponseWriter, r *http.Request, modeler *tim doc, cacheStatus, _, err = QueryCache(ctx, cache, key, nil, modeler.CacheUnmarshaler) if cacheStatus == status.LookupStatusKeyMiss && errors.Is(err, tc.ErrKNF) { cts, doc, elapsed, failedExts, severeFault = fetchTimeseries(pr, trq, client, modeler) + if rlo.FallbackToProxyOnError && len(failedExts) > 0 { + return &dpcResult{cacheStatus: status.LookupStatusProxyOnly}, nil + } if len(failedExts) > 0 && severeFault { return buildErrorResult(doc.StatusCode, doc.SafeHeaderClone(), doc.Body, failedExts), nil } @@ -367,6 +397,9 @@ func DeltaProxyCacheRequest(w http.ResponseWriter, r *http.Request, modeler *tim logging.Pairs{keys.Key: key, keys.BackendName: client.Name(), keys.Detail: err.Error()}) goWithRecover("dpc.cache.Remove.unmarshal", func() { cache.Remove(key) }) cts, doc, elapsed, failedExts, severeFault = fetchTimeseries(pr, trq, client, modeler) + if rlo.FallbackToProxyOnError && len(failedExts) > 0 { + return &dpcResult{cacheStatus: status.LookupStatusProxyOnly}, nil + } if len(failedExts) > 0 && severeFault { return buildErrorResult(doc.StatusCode, doc.SafeHeaderClone(), doc.Body, failedExts), nil } @@ -438,6 +471,9 @@ func DeltaProxyCacheRequest(w http.ResponseWriter, r *http.Request, modeler *tim fetchHeaders := http.Header(doc.Headers).Clone() mts, _, mresp, failedExts, severeFault = fetchExtents(missRanges, frsc, fetchHeaders, client, pr, modeler.WireUnmarshalerReader, span) + if rlo.FallbackToProxyOnError && len(failedExts) > 0 { + return &dpcResult{cacheStatus: status.LookupStatusProxyOnly}, nil + } if len(failedExts) > 0 && severeFault { // mresp.Body is only set inside fetchExtents's non-200 // branch; when every shard fails at the transport level @@ -498,6 +534,10 @@ func DeltaProxyCacheRequest(w http.ResponseWriter, r *http.Request, modeler *tim rts = cts.Clone() } rts.SetTimeRangeQuery(trq) + wireBody, err := prepareDPCResponse(rts, rlo, modeler, doc.StatusCode) + if err != nil { + return &dpcResult{cacheStatus: status.LookupStatusProxyOnly}, nil + } // Crop the Cache Object down to the Sample Size or Age Retention Policy and the // Backfill Tolerance before storing to cache @@ -535,8 +575,7 @@ func DeltaProxyCacheRequest(w http.ResponseWriter, r *http.Request, modeler *tim // provider renders per request (see MarshalVariesByRequest), in // which case each caller marshals the shared timeseries itself rts.SetExtents(nil) // so they are not included in the client response json - var wireBody []byte - if !marshalVaries { + if wireBody == nil && !marshalVaries { var buf bytes.Buffer modeler.WireMarshalWriter(rts, rlo, doc.StatusCode, &buf) wireBody = buf.Bytes() @@ -565,7 +604,7 @@ func DeltaProxyCacheRequest(w http.ResponseWriter, r *http.Request, modeler *tim // handle sentinel statuses that require special responses if result.cacheStatus == status.LookupStatusProxyOnly { - // LRU eviction determined the request is too old to cache + // Retention or provider response validation requires the original query. if trq.OriginalBody != nil { request.SetBody(r, trq.OriginalBody) } @@ -626,6 +665,13 @@ func DeltaProxyCacheRequest(w http.ResponseWriter, r *http.Request, modeler *tim var severeFault bool cts, doc, elapsed, failedExts, severeFault = fetchTimeseries(pr, trq, client, modeler) + if rlo.FallbackToProxyOnError && len(failedExts) > 0 { + if trq.OriginalBody != nil { + request.SetBody(r, trq.OriginalBody) + } + DoProxy(w, r, true) + return + } if len(failedExts) > 0 && severeFault { h := doc.SafeHeaderClone() sc := dpcProxyErrorStatusCode(doc.StatusCode) @@ -636,6 +682,14 @@ func DeltaProxyCacheRequest(w http.ResponseWriter, r *http.Request, modeler *tim } rts = cts.Clone() rts.SetTimeRangeQuery(trq) + wireBody, err := prepareDPCResponse(rts, rlo, modeler, doc.StatusCode) + if err != nil { + if trq.OriginalBody != nil { + request.SetBody(r, trq.OriginalBody) + } + DoProxy(w, r, true) + return + } tspan.SetAttributes(rsc.Tracer, span, attribute.String("cache.status", cacheStatus.String())) @@ -645,10 +699,13 @@ func DeltaProxyCacheRequest(w http.ResponseWriter, r *http.Request, modeler *tim rts.SetExtents(nil) // so they are not included in the client response json rh := doc.SafeHeaderClone() sc := doc.StatusCode + if rsc.TSTransformer != nil { + wireBody = nil + } finalizeDPCResponse(w, r, rsc, rts, rh, sc, cacheStatus, ffStatus, elapsed.Seconds(), missRanges, failedExts, uncachedValueCount, - key, o, rlo, modeler, nil) + key, o, rlo, modeler, wireBody) } func logDeltaRoutine(p logging.Pairs) { @@ -710,9 +767,11 @@ func fetchTimeseries( elapsed = time.Since(start) } + // A fallback may reuse and mutate the request after this function returns. + method, target, userAgent := pr.Method, pr.URL.String(), pr.UserAgent() goWithRecover("dpc.logUpstreamRequest", func() { logUpstreamRequest(o.Name, o.Provider, handlerName, - pr.Method, pr.URL.String(), pr.UserAgent(), resp.StatusCode, 0, elapsed.Seconds()) + method, target, userAgent, resp.StatusCode, 0, elapsed.Seconds()) }) d := &HTTPDocument{ @@ -800,6 +859,7 @@ func fetchExtents( errTs := make(timeseries.ExtentList, len(el)) // the meta-response aggregating all upstream responses mresp := &http.Response{Header: h} + var errorHeaders http.Header // limit concurrent upstream requests to avoid overwhelming the origin eg := errgroup.Group{} @@ -885,14 +945,18 @@ func fetchExtents( var s string if resp.Body != nil { var readErr error - b, readErr = io.ReadAll(io.LimitReader(resp.Body, errorBodyCap)) + b, readErr = io.ReadAll(io.LimitReader(getDecoderReader(resp), errorBodyCap)) if readErr != nil { logger.Warn("failed to read upstream error response body", logging.Pairs{keys.Detail: readErr.Error()}) } s = string(b) respLock.Lock() - mresp.Body = io.NopCloser(bytes.NewReader(b)) + if resp.StatusCode == mresp.StatusCode { + mresp.Body = io.NopCloser(bytes.NewReader(b)) + errorHeaders = resp.Header.Clone() + errorHeaders.Del(headers.NameContentLength) + } respLock.Unlock() } if len(s) > 128 { @@ -921,6 +985,9 @@ func fetchExtents( trimmedList := errTs.TrimEmptyExtents() if trimmedList.Len() == el.Len() { fullFaults = true + if errorHeaders != nil { + mresp.Header = errorHeaders + } } return mts, uncachedValueCount.Load(), mresp, trimmedList, fullFaults diff --git a/pkg/proxy/engines/dpc_response_test.go b/pkg/proxy/engines/dpc_response_test.go new file mode 100644 index 000000000..24813bb8e --- /dev/null +++ b/pkg/proxy/engines/dpc_response_test.go @@ -0,0 +1,80 @@ +/* + * Copyright 2026 The Trickster Authors + * 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 engines + +import ( + "errors" + "io" + "reflect" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/timeseries" + "github.com/trickstercache/trickster/v2/pkg/timeseries/dataset" +) + +func TestPrepareDPCResponse(t *testing.T) { + failure := errors.New("injected serialization failure") + for _, tc := range []struct { + name string + opts timeseries.RequestOptions + body string + err error + calls int + reuse bool + }{ + {"not_fallback", timeseries.RequestOptions{}, "", nil, 0, false}, + {"fast_forward", timeseries.RequestOptions{FallbackToProxyOnError: true}, "result", nil, 1, false}, + {"request_specific", timeseries.RequestOptions{FallbackToProxyOnError: true, FastForwardDisable: true, MarshalVariesByRequest: true}, "result", nil, 1, false}, + {"reuse", timeseries.RequestOptions{FallbackToProxyOnError: true, FastForwardDisable: true}, "result", nil, 1, true}, + {"empty", timeseries.RequestOptions{FallbackToProxyOnError: true, FastForwardDisable: true}, "", nil, 1, true}, + {"failed", timeseries.RequestOptions{FallbackToProxyOnError: true, FastForwardDisable: true}, "partial output", failure, 1, false}, + } { + t.Run(tc.name, func(t *testing.T) { + extents := timeseries.ExtentList{{Start: time.Unix(1, 0), End: time.Unix(10, 0)}} + ts := &dataset.DataSet{ExtentList: extents} + final := ts.Clone() + final.SetExtents(nil) + calls := 0 + modeler := ×eries.Modeler{WireMarshalWriter: func(got timeseries.Timeseries, ro *timeseries.RequestOptions, status int, w io.Writer) error { + calls++ + if got != ts || ro != &tc.opts || status != 200 { + t.Fatal("marshal inputs changed") + } + buffered := tc.opts.FastForwardDisable && !tc.opts.MarshalVariesByRequest + wantExtents := extents + if buffered { + wantExtents = final.Extents() + } + if !reflect.DeepEqual(got.Extents(), wantExtents) { + t.Fatal("pre-render must match final wire extent visibility") + } + if _, err := io.WriteString(w, tc.body); err != nil { + return err + } + return tc.err + }} + body, err := prepareDPCResponse(ts, &tc.opts, modeler, 200) + if !errors.Is(err, tc.err) || calls != tc.calls || (body != nil) != tc.reuse { + t.Fatalf("body=%q err=%v calls=%d", body, err, calls) + } + if tc.reuse && string(body) != tc.body { + t.Fatalf("body=%q want=%q", body, tc.body) + } + if !reflect.DeepEqual(ts.Extents(), extents) { + t.Fatal("cache extents were modified") + } + }) + } +} diff --git a/pkg/proxy/listener/native/native.go b/pkg/proxy/listener/native/native.go index dd75dfc61..ac24b5e04 100644 --- a/pkg/proxy/listener/native/native.go +++ b/pkg/proxy/listener/native/native.go @@ -20,6 +20,7 @@ package native import ( "net/http" "slices" + "strings" "github.com/trickstercache/trickster/v2/pkg/backends" bo "github.com/trickstercache/trickster/v2/pkg/backends/options" @@ -55,7 +56,7 @@ type Adapter interface { ServesProvider(provider string) bool // Providers returns the served provider names, for messages and docs only. Providers() []string - SupportsHTTP() bool + SupportsHTTP(provider string) bool Configured(*listenerconfig.Options) bool ValidateListener(*listenerconfig.Options) error ValidateBackend(*bo.Options) error @@ -76,17 +77,44 @@ func (r Registry) Get(protocol string) Adapter { return r[protocol] } -// GetByProvider returns the adapter whose protocol serves the given backend -// provider, or nil when no native protocol serves it. +// GetByProvider returns the unique adapter serving provider, or nil when the +// provider is unsupported or ambiguous. Use GetForProvider when a listener's +// protocol is known. func (r Registry) GetByProvider(provider string) Adapter { + var found Adapter for _, adapter := range r { if adapter.ServesProvider(provider) { - return adapter + if found != nil { + return nil + } + found = adapter } } + return found +} + +// GetForProvider returns the adapter serving provider over protocol. +func (r Registry) GetForProvider(protocol, provider string) Adapter { + if a := r.Get(protocol); a != nil && a.ServesProvider(provider) { + return a + } return nil } +// ForProvider returns all matching adapters ordered by protocol. +func (r Registry) ForProvider(provider string) []Adapter { + var out []Adapter + for _, a := range r { + if a.ServesProvider(provider) { + out = append(out, a) + } + } + slices.SortFunc(out, func(a, b Adapter) int { + return strings.Compare(a.Protocol(), b.Protocol()) + }) + return out +} + // ConfiguredProtocol returns the first, lexically ordered protocol whose // provider-specific options are present, excluding except. func (r Registry) ConfiguredProtocol(options *listenerconfig.Options, except string) string { diff --git a/pkg/proxy/listener/native/native_test.go b/pkg/proxy/listener/native/native_test.go new file mode 100644 index 000000000..793cc30ab --- /dev/null +++ b/pkg/proxy/listener/native/native_test.go @@ -0,0 +1,46 @@ +/* + * Copyright 2026 The Trickster Authors + * 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 native + +import "testing" + +type testAdapter struct { + Adapter + protocol string +} + +func (a *testAdapter) Protocol() string { return a.protocol } +func (*testAdapter) ServesProvider(name string) bool { return name == "shared" } + +func TestProtocolQualifiedProvider(t *testing.T) { + pg, my := &testAdapter{protocol: "postgres"}, &testAdapter{protocol: "mysql"} + r := Registry{"postgres": pg, "mysql": my} + if r.GetByProvider("shared") != nil { + t.Fatal("ambiguous provider lookup must not depend on map iteration order") + } + if r.GetForProvider("postgres", "shared") != pg || r.GetForProvider("mysql", "shared") != my { + t.Fatal("protocol-qualified lookup did not preserve both adapters") + } + if r.GetForProvider("http", "shared") != nil || r.GetForProvider("mysql", "unknown") != nil { + t.Fatal("unsupported protocol/provider was accepted") + } + all := r.ForProvider("shared") + if len(all) != 2 || all[0] != my || all[1] != pg || len(r.ForProvider("unknown")) != 0 { + t.Fatal("provider adapters must be sorted and complete") + } + delete(r, "mysql") + if r.GetByProvider("shared") != pg { + t.Fatal("unique provider lookup changed") + } +} diff --git a/pkg/proxy/pgwire/adapter.go b/pkg/proxy/pgwire/adapter.go index d34b07343..f795762ee 100644 --- a/pkg/proxy/pgwire/adapter.go +++ b/pkg/proxy/pgwire/adapter.go @@ -19,6 +19,7 @@ import ( "errors" "fmt" "maps" + "net/url" "slices" "strconv" "strings" @@ -50,7 +51,9 @@ func NewNativeListenerAdapter(engines Engines) native.Adapter { return &nativeListenerAdapter{engines: engines} } -func (nativeListenerAdapter) SupportsHTTP() bool { return false } +func (a nativeListenerAdapter) SupportsHTTP(provider string) bool { + return supportsHTTP(a.engines.Get(provider)) +} func (nativeListenerAdapter) Protocol() string { return listenerconfig.ProtocolPostgres } @@ -75,11 +78,21 @@ func (nativeListenerAdapter) ValidateListener(o *listenerconfig.Options) error { } func (a nativeListenerAdapter) ValidateBackend(o *bo.Options) error { + if o != nil && o.HasHTTPListener && a.SupportsHTTP(strings.ToLower(o.Provider)) { + u, err := url.Parse(o.OriginURL) + if err != nil || u.Hostname() == "" || (u.Scheme != "http" && u.Scheme != "https") { + return errors.New("an HTTP listener requires an http(s) origin_url; use a postgres listener for a pgwire-only backend") + } + } if o != nil && o.Postgres != nil { if err := o.Postgres.Validate(); err != nil { return err } } + if o != nil && a.SupportsHTTP(strings.ToLower(o.Provider)) && + !slices.Contains(o.NativeListenerProtocols, listenerconfig.ProtocolPostgres) { + return nil + } _, err := a.configFromOptions(o) return err } diff --git a/pkg/proxy/pgwire/cache.go b/pkg/proxy/pgwire/cache.go index db6c80eed..cdb998a5d 100644 --- a/pkg/proxy/pgwire/cache.go +++ b/pkg/proxy/pgwire/cache.go @@ -168,7 +168,7 @@ func (p *deltaPlan) rowReader(s *session, rowDescription []byte) (*rowReader, er return nil, nativedelta.Unmergeable(errTimeColumn) } decoder, err := newTimeAxisDecoder(kind, p.plan.OutputUnit, - engine.TimeSemantics().NaiveTimestampsAreUTC, s.tracker.setting) + engine.TimeSemantics(), s.tracker.setting) if err != nil { return nil, nativedelta.Unmergeable(err) } diff --git a/pkg/proxy/pgwire/config.go b/pkg/proxy/pgwire/config.go index ca46fdddd..a0afe7120 100644 --- a/pkg/proxy/pgwire/config.go +++ b/pkg/proxy/pgwire/config.go @@ -52,10 +52,10 @@ type Upstream struct { Address string // Host is the origin's host name, used for certificate verification. Host string - // User and Password are the origin credentials from the origin_url. + // User and Password are the credentials from the resolved pgwire URL. User string Password string - // Database is the default database from the origin_url path. + // Database is the default database from the resolved pgwire URL path. Database string // TLS is nil when the upstream TLS mode is disable. TLS *tls.Config @@ -116,7 +116,9 @@ type Config struct { func (c *Config) Terminated() bool { return c.Users != nil } // ConfigFromOptions derives the protocol configuration from backend options. The -// origin URL format is postgres://[user[:password]@]host[:port][/database]. +// pgwire URL format is postgres://[user[:password]@]host[:port][/database]. +// postgres.upstream_url takes precedence over origin_url. HTTP-capable engines +// can also use the HTTP origin's host with their default native port. func ConfigFromOptions(o *bo.Options, engine Engine) (Config, error) { if o == nil { return Config{}, errors.New("nil postgres backend options") @@ -185,12 +187,26 @@ func (c *Config) ApplyListenerOptions(o *pgo.ListenerOptions) { } func upstreamFromOptions(o *bo.Options, engine Engine) (Upstream, error) { - u, err := url.Parse(o.OriginURL) + rawURL := o.OriginURL + overridden := o.Postgres != nil && o.Postgres.UpstreamURL != "" + if overridden { + rawURL = o.Postgres.UpstreamURL + } + u, err := url.Parse(rawURL) if err != nil { - return Upstream{}, fmt.Errorf("parse postgres origin URL: %w", err) + return Upstream{}, errors.New("parse postgres origin URL: invalid URL") } if u.Scheme != schemePostgres && u.Scheme != schemePostgreSQL { - return Upstream{}, fmt.Errorf("unsupported postgres origin scheme %q", u.Scheme) + if overridden || !supportsHTTP(engine) || (u.Scheme != "http" && u.Scheme != "https") { + return Upstream{}, fmt.Errorf("unsupported postgres origin scheme %q", u.Scheme) + } + // HTTP userinfo, port, path and query belong to another protocol. + // Only the host can be shared without an explicit pgwire URL. + host := u.Hostname() + if host == "" { + return Upstream{}, errors.New("postgres origin URL has no host") + } + u = &url.URL{Scheme: schemePostgres, Host: net.JoinHostPort(host, engine.DefaultPort())} } host := u.Hostname() if host == "" { diff --git a/pkg/proxy/pgwire/config_test.go b/pkg/proxy/pgwire/config_test.go index d25f1b689..220505c7c 100644 --- a/pkg/proxy/pgwire/config_test.go +++ b/pkg/proxy/pgwire/config_test.go @@ -63,6 +63,7 @@ func (e testEngine) Analyzer() sqlanalyzer.DialectAnalyzer { } return testAnalyzer } + func (e testEngine) Defaults() EngineDefaults { return EngineDefaults{UpstreamTLSMode: e.tlsMode} } func (testEngine) TimeAxis(oid uint32) (TimeAxisKind, bool) { return StandardTimeAxis(oid) } func (testEngine) TimeSemantics() TimeSemantics { return TimeSemantics{} } @@ -302,7 +303,7 @@ func adapterTestConfig() *config.Config { func TestNativeListenerAdapterContract(t *testing.T) { adapter := NewNativeListenerAdapter(NewEngines(testEngine{})) - if adapter.Protocol() != listenerconfig.ProtocolPostgres || adapter.SupportsHTTP() { + if adapter.Protocol() != listenerconfig.ProtocolPostgres || adapter.SupportsHTTP(providers.Postgres) { t.Fatalf("unexpected identity %q", adapter.Protocol()) } if !adapter.ServesProvider(providers.Postgres) || !adapter.ServesProvider(providers.TimescaleDB) || diff --git a/pkg/proxy/pgwire/engine.go b/pkg/proxy/pgwire/engine.go index 69f83a260..b3762615d 100644 --- a/pkg/proxy/pgwire/engine.go +++ b/pkg/proxy/pgwire/engine.go @@ -72,6 +72,50 @@ type TimeSemantics struct { // NaiveTimestampsAreUTC means TIMESTAMP WITHOUT TIME ZONE values are UTC // instants whatever the session time zone is. PostgreSQL gives them no zone. NaiveTimestampsAreUTC bool + // LosslessFloatText means float text always round-trips, independently of + // PostgreSQL's extra_float_digits setting. Otherwise that setting is required. + LosslessFloatText bool + // Assumed values are engine guarantees, used only when a setting is unknown. + // They must not stand in for configurable server or role defaults. + AssumedDateStyle string + AssumedTimeZone string + AssumedIntervalStyle string + AssumedStandardConformingStrings string + AssumedIntegerDatetimes string +} + +// SessionDefaultsProbe reads effective settings after an origin login. Each +// result must contain one non-null row; Names follows the flattened column order. +// An empty SQL string disables the probe without supplying any settings. +type SessionDefaultsProbe struct { + SQL string + Names []string +} + +// SessionDefaultsEngine overrides PostgreSQL's session-defaults query. +type SessionDefaultsEngine interface { + SessionDefaultsProbe() SessionDefaultsProbe +} + +// SessionSettings describes settings followed from successful client statements. +// Names and alias targets are lower case. Tables must be immutable after use. +type SessionSettings struct { + // Tracked replaces PostgreSQL's client-tracked table. Listed settings are + // followed even if the origin announced an initial value at login. + Tracked map[string]struct{} + // Neutral replaces PostgreSQL's table of settings that do not shape results. + Neutral map[string]struct{} + Aliases map[string]string + // UnconfirmedStartup partitions by requested startup parameters without + // treating them as effective settings until announced or probed. + UnconfirmedStartup bool + // LocalPersists is for origins whose SET LOCAL is session-scoped. + LocalPersists bool +} + +// SessionSettingsEngine overrides PostgreSQL's client-setting semantics. +type SessionSettingsEngine interface { + SessionSettings() SessionSettings } // EngineDefaults are the connection settings assumed when a backend sets none. @@ -111,6 +155,17 @@ type Engine interface { TimeSemantics() TimeSemantics } +// HTTPEngine is an optional capability for providers that also expose HTTP. +// Engines without it retain PostgreSQL's native-only behavior. +type HTTPEngine interface { + SupportsHTTP() bool +} + +func supportsHTTP(engine Engine) bool { + httpEngine, ok := engine.(HTTPEngine) + return ok && httpEngine.SupportsHTTP() +} + // Engines is the explicit registry of engines, keyed by canonical provider name. type Engines map[string]Engine diff --git a/pkg/proxy/pgwire/gate_test.go b/pkg/proxy/pgwire/gate_test.go index 24d343234..2c8c82eec 100644 --- a/pkg/proxy/pgwire/gate_test.go +++ b/pkg/proxy/pgwire/gate_test.go @@ -568,7 +568,7 @@ func TestUnannouncedSettingsComeFromTheSessionsDefaults(t *testing.T) { // a role or database default is invisible on the wire, so it is read once at origin login s := gateTestSession(t, nil) floatAxis := func() error { - _, err := newTimeAxisDecoder(TimeAxisEpochFloat, timeseries.DateTimeUnixSecs, false, s.tracker.setting) + _, err := newTimeAxisDecoder(TimeAxisEpochFloat, timeseries.DateTimeUnixSecs, TimeSemantics{}, s.tracker.setting) return err } apply := func(sql string) { diff --git a/pkg/proxy/pgwire/options/options.go b/pkg/proxy/pgwire/options/options.go index 8810a9548..3bd8011b6 100644 --- a/pkg/proxy/pgwire/options/options.go +++ b/pkg/proxy/pgwire/options/options.go @@ -66,6 +66,9 @@ func TLSModes() []string { return slices.Clone(tlsModes) } // Options contains settings for a backend reached over the PostgreSQL wire protocol. type Options struct { + // UpstreamURL overrides origin_url for the PostgreSQL wire protocol. + // It includes the native host, optional credentials and database. + UpstreamURL string `yaml:"upstream_url,omitempty"` // UpstreamTLSMode selects TLS toward the origin, independent of listener // TLS. Empty selects the engine's default, which is disable for PostgreSQL. UpstreamTLSMode string `yaml:"upstream_tls_mode,omitempty"` diff --git a/pkg/proxy/pgwire/result_test.go b/pkg/proxy/pgwire/result_test.go index 1fb84cc44..486ee2ff2 100644 --- a/pkg/proxy/pgwire/result_test.go +++ b/pkg/proxy/pgwire/result_test.go @@ -242,7 +242,7 @@ func TestTimeAxisDecoding(t *testing.T) { "float epoch": {TimeAxisEpochFloat, timeseries.DateTimeUnixSecs, "1789027500", want}, "numeric epoch": {TimeAxisEpochNumeric, timeseries.DateTimeUnixSecs, "1789027500.000", want}, } { - decoder, err := newTimeAxisDecoder(test.kind, test.unit, false, iso) + decoder, err := newTimeAxisDecoder(test.kind, test.unit, TimeSemantics{}, iso) if err != nil { t.Fatalf("%s: %v", name, err) } @@ -250,7 +250,7 @@ func TestTimeAxisDecoding(t *testing.T) { t.Fatalf("%s: got %v, %v; want %v", name, got, err, test.want) } } - zoned, _ := newTimeAxisDecoder(TimeAxisTimestampTZ, 0, false, iso) + zoned, _ := newTimeAxisDecoder(TimeAxisTimestampTZ, 0, TimeSemantics{}, iso) for _, text := range []string{ "", "infinity", "-infinity", "12026-09-10 08:05:00+00", "2026-09-10 08:05:00+00 BC", "2026-09-10 08:05:00", "2026-09-10 08:05:00Z", "2026-09-10 08:05:00.+00", "2026-09-10 08:05:00.1234567890+00", "2026-13-10 08:05:00+00", @@ -261,21 +261,21 @@ func TestTimeAxisDecoding(t *testing.T) { t.Fatalf("%q must fail closed, got %v", text, err) } } - naive, _ := newTimeAxisDecoder(TimeAxisTimestamp, 0, false, iso) + naive, _ := newTimeAxisDecoder(TimeAxisTimestamp, 0, TimeSemantics{}, iso) if _, err := naive.decode([]byte("2026-09-10 08:05:00+00")); !errors.Is(err, errTimeAxis) { t.Fatalf("a zone-less column cannot carry an offset, got %v", err) } - date, _ := newTimeAxisDecoder(TimeAxisDate, 0, false, iso) + date, _ := newTimeAxisDecoder(TimeAxisDate, 0, TimeSemantics{}, iso) if _, err := date.decode([]byte("2026-09-10 BC")); !errors.Is(err, errTimeAxis) { t.Fatalf("expected a BC date to fail closed, got %v", err) } - epoch, _ := newTimeAxisDecoder(TimeAxisEpochFloat, timeseries.DateTimeUnixSecs, false, iso) + epoch, _ := newTimeAxisDecoder(TimeAxisEpochFloat, timeseries.DateTimeUnixSecs, TimeSemantics{}, iso) for _, text := range []string{"2e+09x", "1789027500.5", "1e300"} { if _, err := epoch.decode([]byte(text)); !errors.Is(err, errTimeAxis) { t.Fatalf("%q must fail closed, got %v", text, err) } } - integer, _ := newTimeAxisDecoder(TimeAxisEpochInteger, timeseries.DateTimeUnixSecs, false, iso) + integer, _ := newTimeAxisDecoder(TimeAxisEpochInteger, timeseries.DateTimeUnixSecs, TimeSemantics{}, iso) for _, text := range []string{"abc", "9223372036854775807"} { if _, err := integer.decode([]byte(text)); !errors.Is(err, errTimeAxis) { t.Fatalf("%q must fail closed, got %v", text, err) @@ -306,7 +306,7 @@ func TestTimeAxisDecoderFailsClosedOnSessionSettings(t *testing.T) { "float with exact digits": {TimeAxisEpochFloat, timeseries.DateTimeUnixNano, false, map[string]string{varExtraFloatDigits: "3"}, true}, "unknown kind": {0, 0, false, nil, false}, } { - _, err := newTimeAxisDecoder(test.kind, test.unit, test.naiveUTC, timeAxisSettings(test.settings)) + _, err := newTimeAxisDecoder(test.kind, test.unit, TimeSemantics{NaiveTimestampsAreUTC: test.naiveUTC}, timeAxisSettings(test.settings)) if (err == nil) != test.ok { t.Errorf("%s: err = %v, want ok = %t", name, err, test.ok) } @@ -314,7 +314,7 @@ func TestTimeAxisDecoderFailsClosedOnSessionSettings(t *testing.T) { } func TestBucketTime(t *testing.T) { - decoder, _ := newTimeAxisDecoder(TimeAxisTimestampTZ, 0, false, timeAxisSettings(map[string]string{"datestyle": "ISO"})) + decoder, _ := newTimeAxisDecoder(TimeAxisTimestampTZ, 0, TimeSemantics{}, timeAxisSettings(map[string]string{"datestyle": "ISO"})) row := func(values ...[]byte) []byte { body, _ := (&pgproto3.DataRow{Values: values}).Encode(nil) return body[frameHeaderLen:] diff --git a/pkg/proxy/pgwire/route.go b/pkg/proxy/pgwire/route.go index f0d537465..c4c14256f 100644 --- a/pkg/proxy/pgwire/route.go +++ b/pkg/proxy/pgwire/route.go @@ -105,7 +105,7 @@ func (s *session) route() bool { metrics.PGWireRouteSelections.WithLabelValues(router, target.config.BackendName, string(decision.Outcome)).Inc() s.server = target if target.config.Analyzer != nil { - s.tracker = newSessionTracker(s.user, s.database, s.params) + s.tracker = newSessionTracker(s.user, s.database, s.params, s.server.config.Engine) } return true } diff --git a/pkg/proxy/pgwire/session.go b/pkg/proxy/pgwire/session.go index f97dcdda4..cd50c15fe 100644 --- a/pkg/proxy/pgwire/session.go +++ b/pkg/proxy/pgwire/session.go @@ -262,7 +262,7 @@ func (s *session) acceptStartupMessage(version uint32, packet []byte) (bool, err } s.database, s.params, s.rawStartup, s.minor = params[paramDatabase], params, packet, version&0xffff if s.server.config.Analyzer != nil { - s.tracker = newSessionTracker(s.user, s.database, params) + s.tracker = newSessionTracker(s.user, s.database, params, s.server.config.Engine) } return true, nil } diff --git a/pkg/proxy/pgwire/session_engine_test.go b/pkg/proxy/pgwire/session_engine_test.go new file mode 100644 index 000000000..df13bd166 --- /dev/null +++ b/pkg/proxy/pgwire/session_engine_test.go @@ -0,0 +1,259 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 pgwire + +import ( + "context" + "errors" + "reflect" + "slices" + "testing" + + "github.com/trickstercache/trickster/v2/pkg/timeseries" + + "github.com/jackc/pgx/v5/pgconn" +) + +type sessionTestEngine struct { + testEngine + probe SessionDefaultsProbe + semantics TimeSemantics + settings SessionSettings +} + +func (e sessionTestEngine) SessionDefaultsProbe() SessionDefaultsProbe { return e.probe } +func (e sessionTestEngine) TimeSemantics() TimeSemantics { return e.semantics } +func (e sessionTestEngine) SessionSettings() SessionSettings { return e.settings } + +func TestEngineSessionProbe(t *testing.T) { + for name, probe := range map[string]SessionDefaultsProbe{ + "none": {}, + "custom": {SQL: "SELECT 1", Names: []string{"sample"}}, + } { + t.Run(name, func(t *testing.T) { + upstream := newFakeUpstream(t, cleartextUpstream) + config := terminatedConfig(upstream, testClientPass) + config.Engine = sessionTestEngine{probe: probe} + conn, defaults, err := config.loginUpstream(t.Context(), "", nil, true) + if err != nil { + t.Fatal(err) + } + _ = conn.Conn.Close() + if slices.Contains(upstream.received(), unannouncedSettingsSQL) { + t.Fatal("used PostgreSQL's probe for an engine override") + } + if probe.SQL == "" { + if len(defaults) != 0 || len(upstream.received()) != 0 { + t.Fatal("an empty probe must not query or invent defaults") + } + } else if defaults["sample"] != "1" || !slices.Contains(upstream.received(), probe.SQL) { + t.Fatalf("custom probe not applied: %v", defaults) + } + }) + } + for _, query := range []string{fakeQueryError, fakeQueryMany, fakeQueryPing} { + t.Run(query, func(t *testing.T) { + upstream := newFakeUpstream(t, cleartextUpstream) + config := terminatedConfig(upstream, testClientPass) + config.Engine = sessionTestEngine{probe: SessionDefaultsProbe{SQL: query, Names: []string{"sample"}}} + ctx, cancel := context.WithTimeout(t.Context(), fakeTimeout) + defer cancel() + conn, _, err := config.loginUpstream(ctx, "", nil, true) + if err == nil { + _ = conn.Conn.Close() + t.Fatal("a failed or malformed probe must fail the login") + } + }) + } +} + +func TestSettingsProbeResults(t *testing.T) { + row := func(values ...string) *pgconn.Result { + r := &pgconn.Result{Rows: [][][]byte{make([][]byte, len(values))}} + for i, v := range values { + r.Rows[0][i] = []byte(v) + } + return r + } + for name, results := range map[string][]*pgconn.Result{ + "one select": {row("UTC", "ISO")}, "two show statements": {row("UTC"), row("ISO")}, + } { + t.Run(name, func(t *testing.T) { + got, err := settingsFromResults(results, []string{"timezone", "datestyle"}) + if err != nil || !reflect.DeepEqual(got, map[string]string{"timezone": "UTC", "datestyle": "ISO"}) { + t.Fatalf("got %v, %v", got, err) + } + }) + } + for name, results := range map[string][]*pgconn.Result{ + "empty": {}, "nil result": {nil}, "no rows": {{}}, "empty row": {row()}, + "too few": {row("UTC")}, "too many": {row("UTC", "ISO", "extra")}, + "null": {{Rows: [][][]byte{{nil, []byte("ISO")}}}}, + "extra row": {{Rows: [][][]byte{{[]byte("UTC"), []byte("ISO")}, {[]byte("UTC"), []byte("ISO")}}}}, + "error": {{Err: errors.New("failed")}}, + } { + t.Run(name, func(t *testing.T) { + if _, err := settingsFromResults(results, []string{"timezone", "datestyle"}); err == nil { + t.Fatal("accepted a malformed settings result") + } + }) + } + for _, names := range [][]string{{}, {"a", "a"}, {"a", ""}} { + if _, err := settingsFromResults([]*pgconn.Result{row("UTC", "ISO")}, names); err == nil { + t.Fatalf("accepted names %v", names) + } + } +} + +func trackerEngine() sessionTestEngine { + return sessionTestEngine{ + semantics: TimeSemantics{AssumedTimeZone: "UTC", AssumedDateStyle: "ISO", AssumedStandardConformingStrings: "on"}, + settings: SessionSettings{ + Tracked: map[string]struct{}{varTimeZone: {}, "datestyle": {}, varStandardConformingStrings: {}}, + Neutral: map[string]struct{}{paramApplicationName: {}}, + Aliases: map[string]string{"time_zone": varTimeZone}, LocalPersists: true, + }, + } +} + +func TestAssumedSessionSettings(t *testing.T) { + engine := trackerEngine() + assumed := newSessionTracker("user", "db", nil, engine) + announced := newSessionTracker("user", "db", nil) + announced.parameterStatus("TimeZone", "UTC") + announced.parameterStatus("DateStyle", "ISO") + announced.parameterStatus(varStandardConformingStrings, "on") + if !assumed.utc() || assumed.lexicalOptions() || assumed.sessionIdentity() != announced.sessionIdentity() { + t.Fatal("assumptions must have the same identity as equal announcements") + } + initial := assumed.sessionIdentity() + assumed.parameterStatus("TimeZone", "Asia/Kolkata") + if assumed.utc() || assumed.sessionIdentity() == initial { + t.Fatal("real announcements must override assumptions") + } + assumed.sessionDefaults(map[string]string{"datestyle": "SQL", varStandardConformingStrings: "off"}) + if style, _ := assumed.setting("datestyle"); style != "SQL" || !assumed.lexicalOptions() { + t.Fatal("probed values must override assumptions") + } + startup := newSessionTracker("user", "db", map[string]string{"time_zone": "Asia/Kolkata"}, engine) + if startup.utc() { + t.Fatal("startup alias must override an assumed zone") + } + if ok, _ := startup.cacheable(); !ok { + t.Fatal("a known startup setting disabled the cache") + } + startup.sessionDefaults(map[string]string{varTimeZone: "UTC"}) + if !startup.utc() { + t.Fatal("the effective probe must override a requested startup value") + } + engine.semantics.AssumedTimeZone = "" + engine.settings.UnconfirmedStartup = true + unconfirmed := newSessionTracker("user", "db", map[string]string{"TimeZone": "UTC"}, engine) + if unconfirmed.utc() { + t.Fatal("an unconfirmed startup parameter established the effective timezone") + } + other := newSessionTracker("user", "db", map[string]string{"TimeZone": "Asia/Kolkata"}, engine) + if other.sessionIdentity() == unconfirmed.sessionIdentity() { + t.Fatal("unconfirmed startup parameters must still partition the cache") + } + unconfirmed.parameterStatus("TimeZone", "UTC") + if !unconfirmed.utc() { + t.Fatal("a real announcement must confirm the startup timezone") + } +} + +func TestEngineClientSettings(t *testing.T) { + for _, sql := range []string{"SET time_zone = 'Asia/Kolkata'", "SET LOCAL time_zone = 'Asia/Kolkata'", "SET TIME ZONE 'Asia/Kolkata'"} { + t.Run(sql, func(t *testing.T) { + tr := newSessionTracker("user", "db", nil, trackerEngine()) + // Some origins announce only the initial value and never updates. + tr.parameterStatus("TimeZone", "UTC") + tr.sessionDefaults(map[string]string{varTimeZone: "UTC"}) + initial := tr.sessionIdentity() + class := classify(sql, false) + tr.observe(&class, true, true, false) + if ok, _ := tr.cacheable(); ok { + t.Fatal("pending SET can be cached") + } + tr.ready(true) + if !tr.utc() || tr.sessionIdentity() != initial { + t.Fatal("failed SET changed the session") + } + tr.observe(&class, true, true, false) + tr.ready(false) + if tr.utc() || tr.sessionIdentity() == initial { + t.Fatal("successful SET was lost") + } + if ok, reason := tr.cacheable(); !ok { + t.Fatalf("successful modeled SET disabled cache: %s", reason) + } + reset := classify("RESET time_zone", false) + tr.observe(&reset, true, true, false) + tr.ready(false) + if !tr.utc() { + t.Fatal("RESET lost the initial probed value") + } + }) + } + tr := newSessionTracker("user", "db", nil, trackerEngine()) + class := classify("SET time_zone = 'GMT'", false) + tr.observe(&class, true, true, false) + tr.parameterStatus("TimeZone", "UTC") + tr.ready(false) + if zone, _ := tr.setting(varTimeZone); zone != "UTC" { + t.Fatal("client text overrode a real announcement") + } + for _, flags := range [][3]bool{{false, true, false}, {true, false, false}, {true, true, true}} { + tr := newSessionTracker("user", "db", nil, trackerEngine()) + tr.observe(&class, flags[0], flags[1], flags[2]) + if ok, _ := tr.cacheable(); ok { + t.Fatalf("uncertain SET context %v was cacheable", flags) + } + } + tr = newSessionTracker("user", "db", nil, trackerEngine()) + class = classify("SET standard_conforming_strings = off", false) + tr.observe(&class, true, true, false) + tr.ready(false) + if !tr.lexicalOptions() { + t.Fatal("client-only string setting did not update the lexer") + } + class = classify("RESET standard_conforming_strings", true) + tr.observe(&class, true, true, false) + tr.ready(false) + if tr.lexicalOptions() { + t.Fatal("RESET did not restore the assumed lexer mode") + } +} + +func TestEngineLosslessFloatText(t *testing.T) { + unknown := timeAxisSettings(nil) + if _, err := newTimeAxisDecoder(TimeAxisEpochFloat, timeseries.DateTimeUnixSecs, TimeSemantics{}, unknown); err == nil { + t.Fatal("an unprobed PostgreSQL float must still fail closed") + } + decoder, err := newTimeAxisDecoder(TimeAxisEpochFloat, timeseries.DateTimeUnixSecs, TimeSemantics{LosslessFloatText: true}, unknown) + if err != nil { + t.Fatal(err) + } + if _, err := decoder.decode([]byte("1700000100.0")); err != nil { + t.Fatal(err) + } + for _, value := range []string{"1700000100.5", "NaN", "Infinity", "9007199254740994"} { + if _, err := decoder.decode([]byte(value)); err == nil { + t.Fatalf("accepted unsafe epoch %s", value) + } + } +} diff --git a/pkg/proxy/pgwire/timeaxis.go b/pkg/proxy/pgwire/timeaxis.go index 6e06e4a19..a5c0ed3eb 100644 --- a/pkg/proxy/pgwire/timeaxis.go +++ b/pkg/proxy/pgwire/timeaxis.go @@ -52,7 +52,7 @@ type timeAxisDecoder struct { unit timeseries.FieldDataType } -func newTimeAxisDecoder(kind TimeAxisKind, unit timeseries.FieldDataType, naiveUTC bool, +func newTimeAxisDecoder(kind TimeAxisKind, unit timeseries.FieldDataType, semantics TimeSemantics, settings func(string) (string, bool), ) (*timeAxisDecoder, error) { // checks that a column of this kind can be decoded under @@ -63,7 +63,7 @@ func newTimeAxisDecoder(kind TimeAxisKind, unit timeseries.FieldDataType, naiveU if style, ok := settings("datestyle"); !ok || !strings.HasPrefix(strings.ToUpper(style), settingISO) { return nil, errTimeAxis } - if kind != TimeAxisTimestampTZ && !naiveUTC { + if kind != TimeAxisTimestampTZ && !semantics.NaiveTimestampsAreUTC { // a zone-less value is compared in the session zone at the origin if zone, ok := settings(varTimeZone); !ok || !isUTCZone(zone) { return nil, errTimeAxis @@ -73,7 +73,7 @@ func newTimeAxisDecoder(kind TimeAxisKind, unit timeseries.FieldDataType, naiveU if !isEpochUnit(unit) { return nil, errTimeAxis } - if kind == TimeAxisEpochFloat { + if kind == TimeAxisEpochFloat && !semantics.LosslessFloatText { // negative extra_float_digits rounds an epoch to text like 2e+09, which can still land // on the grid. The origin never announces the setting, so an unknown value fails closed. digits, ok := settings(varExtraFloatDigits) diff --git a/pkg/proxy/pgwire/tracker.go b/pkg/proxy/pgwire/tracker.go index 909e6d548..54c667540 100644 --- a/pkg/proxy/pgwire/tracker.go +++ b/pkg/proxy/pgwire/tracker.go @@ -17,6 +17,7 @@ package pgwire import ( "encoding/binary" + "maps" "slices" "strings" "sync" @@ -73,29 +74,47 @@ type sessionTracker struct { // defaults holds what the session started with for settings the origin never // announces; a RESET returns to them. Empty when the origin login is the client's own. defaults map[string]string + assumed map[string]string + settings SessionSettings pending *statementClass + pendingReported bool unsafe string identity string backslashEscapes bool } -func newSessionTracker(user, database string, params map[string]string) *sessionTracker { +func newSessionTracker(user, database string, params map[string]string, engines ...Engine) *sessionTracker { t := &sessionTracker{ user: user, database: database, reported: make(map[string]string), client: make(map[string]string), + settings: SessionSettings{Tracked: clientIdentity, Neutral: neutralSettings}, + } + if len(engines) > 0 && engines[0] != nil { + engine := engines[0] + if settings, ok := engine.(SessionSettingsEngine); ok { + t.settings = settings.SessionSettings() + } + semantics := engine.TimeSemantics() + t.assumed = map[string]string{ + "datestyle": semantics.AssumedDateStyle, varTimeZone: semantics.AssumedTimeZone, + "intervalstyle": semantics.AssumedIntervalStyle, + varStandardConformingStrings: semantics.AssumedStandardConformingStrings, + "integer_datetimes": semantics.AssumedIntegerDatetimes, + } + maps.DeleteFunc(t.assumed, func(_, value string) bool { return value == "" }) } for name, value := range params { - name = strings.ToLower(name) + name = t.settingName(name) switch { case name == paramUser || name == paramDatabase || name == paramReplication || strings.HasPrefix(name, protocolOptionPrefix): case name == paramOptions: t.options = value default: - if _, neutral := neutralSettings[name]; neutral { + if _, neutral := t.settings.Neutral[name]; neutral { continue } - if !t.modeled(name) { + if !t.modeled(name) || t.settings.UnconfirmedStartup { // an unknown startup setting is constant for the session, so it // partitions the cache instead of disabling it name = startupSettingPrefix + name @@ -103,19 +122,28 @@ func newSessionTracker(user, database string, params map[string]string) *session t.client[name] = value } } + t.updateLexicalOptions() return t } +func (t *sessionTracker) settingName(name string) string { + name = strings.ToLower(name) + if canonical, ok := t.settings.Aliases[name]; ok { + return canonical + } + return name +} + func (t *sessionTracker) modeled(name string) bool { _, reported := reportedIdentity[name] - _, client := clientIdentity[name] + _, client := t.settings.Tracked[name] return reported || client } func (t *sessionTracker) parameterStatus(name, value string) { // records a setting the origin announced. It is authoritative: // it reflects SET, RESET, rollbacks and function side effects alike. - name = strings.ToLower(name) + name = t.settingName(name) if _, ok := reportedIdentity[name]; !ok { return } @@ -124,9 +152,10 @@ func (t *sessionTracker) parameterStatus(name, value string) { t.reported[name] = value delete(t.client, name) t.identity = "" - if name == varStandardConformingStrings { - t.backslashEscapes = value == settingOff + if t.pending != nil && t.pending.name == name { + t.pendingReported = true } + t.updateLexicalOptions() } func (t *sessionTracker) lexicalOptions() bool { @@ -144,18 +173,21 @@ func (t *sessionTracker) observe(class *statementClass, txIdle, settled, extende t.markUnsafe(unsafeStatement) return } - if class.kind != stmtSet && class.kind != stmtReset && class.kind != stmtDiscardAll || class.local { + if class.kind != stmtSet && class.kind != stmtReset && class.kind != stmtDiscardAll || + class.local && !t.settings.LocalPersists { return } + class.name = t.settingName(class.name) if class.kind != stmtDiscardAll && class.name != varAll { - if _, neutral := neutralSettings[class.name]; neutral { + if _, neutral := t.settings.Neutral[class.name]; neutral { return } if !t.modeled(class.name) { t.markUnsafe(unsafeSetting) return } - if _, announced := t.reported[class.name]; announced { + _, tracked := t.settings.Tracked[class.name] + if _, announced := t.reported[class.name]; announced && !tracked { return } } @@ -170,6 +202,7 @@ func (t *sessionTracker) observe(class *statementClass, txIdle, settled, extende t.markUnsafe(unsafeSetInTx) default: t.pending = class + t.pendingReported = false } } @@ -181,7 +214,7 @@ func (t *sessionTracker) ready(failed bool) { return } t.pending = nil - if failed { + if failed || t.pendingReported { return } t.identity = "" @@ -194,9 +227,12 @@ func (t *sessionTracker) ready(failed bool) { } case class.isDefault: delete(t.client, class.name) + delete(t.reported, class.name) default: t.client[class.name] = class.value + delete(t.reported, class.name) } + t.updateLexicalOptions() } func (t *sessionTracker) setting(name string) (string, bool) { @@ -204,20 +240,38 @@ func (t *sessionTracker) setting(name string) (string, bool) { // the origin announced over what the client asked for. t.mtx.Lock() defer t.mtx.Unlock() + return t.settingLocked(name) +} + +func (t *sessionTracker) settingLocked(name string) (string, bool) { if value, ok := t.reported[name]; ok { return value, true } if value, ok := t.client[name]; ok { return value, true } - value, ok := t.defaults[name] + if value, ok := t.defaults[name]; ok { + return value, true + } + value, ok := t.assumed[name] return value, ok } +func (t *sessionTracker) updateLexicalOptions() { + value, _ := t.settingLocked(varStandardConformingStrings) + t.backslashEscapes = strings.EqualFold(value, settingOff) +} + func (t *sessionTracker) sessionDefaults(defaults map[string]string) { t.mtx.Lock() defer t.mtx.Unlock() t.defaults, t.identity = defaults, "" + // The probe observes the effective startup state, even if the origin ignored + // a requested parameter. Do not let the request override that observation. + for name := range defaults { + delete(t.client, name) + } + t.updateLexicalOptions() } func (t *sessionTracker) utc() bool { @@ -253,7 +307,19 @@ func (t *sessionTracker) sessionIdentity() string { appendIdentityField(&identity, t.user) appendIdentityField(&identity, t.database) appendIdentityField(&identity, t.options) - for _, settings := range []map[string]string{t.reported, t.client, t.defaults} { + reported := t.reported + if len(t.assumed) > 0 { + reported = maps.Clone(t.reported) + for name, value := range t.assumed { + _, announced := reported[name] + _, client := t.client[name] + _, probed := t.defaults[name] + if !announced && !client && !probed { + reported[name] = value + } + } + } + for _, settings := range []map[string]string{reported, t.client, t.defaults} { names := make([]string, 0, len(settings)) for name := range settings { names = append(names, name) diff --git a/pkg/proxy/pgwire/upstream.go b/pkg/proxy/pgwire/upstream.go index 126ea7eea..18b082237 100644 --- a/pkg/proxy/pgwire/upstream.go +++ b/pkg/proxy/pgwire/upstream.go @@ -142,7 +142,11 @@ func (c *Config) loginUpstream(ctx context.Context, database string, } var defaults map[string]string if probe { - if defaults, err = unannouncedSettings(ctx, conn); err != nil { + settingsProbe := SessionDefaultsProbe{SQL: unannouncedSettingsSQL, Names: unannouncedSettingNames} + if engine, ok := c.Engine.(SessionDefaultsEngine); ok { + settingsProbe = engine.SessionDefaultsProbe() + } + if defaults, err = sessionDefaults(ctx, conn, settingsProbe); err != nil { _ = conn.Close(ctx) return nil, nil, fmt.Errorf("postgres upstream login: %w", sanitizeConnectError(err)) } @@ -155,17 +159,37 @@ func (c *Config) loginUpstream(ctx context.Context, database string, return hijacked, defaults, nil } -func unannouncedSettings(ctx context.Context, conn *pgconn.PgConn) (map[string]string, error) { - results, err := conn.Exec(ctx, unannouncedSettingsSQL).ReadAll() +func sessionDefaults(ctx context.Context, conn *pgconn.PgConn, probe SessionDefaultsProbe) (map[string]string, error) { + if probe.SQL == "" { + return nil, nil + } + results, err := conn.Exec(ctx, probe.SQL).ReadAll() if err != nil { return nil, err } - if len(results) != 1 || len(results[0].Rows) != 1 || len(results[0].Rows[0]) != len(unannouncedSettingNames) { - return nil, errors.New("unexpected answer to the settings probe") + return settingsFromResults(results, probe.Names) +} + +func settingsFromResults(results []*pgconn.Result, names []string) (map[string]string, error) { + defaults := make(map[string]string, len(names)) + at := 0 + for _, result := range results { + if result == nil || result.Err != nil || len(result.Rows) != 1 || len(result.Rows[0]) == 0 { + return nil, errors.New("unexpected answer to the settings probe") + } + for _, value := range result.Rows[0] { + if at >= len(names) || value == nil || names[at] == "" { + return nil, errors.New("unexpected answer to the settings probe") + } + if _, duplicate := defaults[names[at]]; duplicate { + return nil, errors.New("duplicate name in the settings probe") + } + defaults[names[at]] = string(value) + at++ + } } - defaults := make(map[string]string, len(unannouncedSettingNames)) - for i, name := range unannouncedSettingNames { - defaults[name] = string(results[0].Rows[0][i]) + if at == 0 || at != len(names) { + return nil, errors.New("unexpected answer to the settings probe") } return defaults, nil } diff --git a/pkg/proxy/pgwire/upstream_url_test.go b/pkg/proxy/pgwire/upstream_url_test.go new file mode 100644 index 000000000..f5a74890f --- /dev/null +++ b/pkg/proxy/pgwire/upstream_url_test.go @@ -0,0 +1,155 @@ +/* + * Copyright 2026 The Trickster Authors + * + * 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 pgwire + +import ( + "strings" + "testing" + + bo "github.com/trickstercache/trickster/v2/pkg/backends/options" + pgo "github.com/trickstercache/trickster/v2/pkg/proxy/pgwire/options" + + "go.yaml.in/yaml/v3" +) + +type mixedTestEngine struct{ testEngine } + +func (mixedTestEngine) Name() string { return "mixed" } +func (mixedTestEngine) DefaultPort() string { return "4003" } +func (mixedTestEngine) SupportsHTTP() bool { return true } + +func TestMixedUpstreamResolution(t *testing.T) { + for _, tt := range []struct { + name, origin, override string + want Upstream + }{ + { + "HTTP host only", "http://http-user:http-secret@db.example:4000/v1/sql?db=other", "", + Upstream{Address: "db.example:4003", Host: "db.example"}, + }, + { + "HTTPS IPv6", "https://[::1]:4000/v1/prometheus", "", + Upstream{Address: "[::1]:4003", Host: "::1"}, + }, + { + "explicit PG override", "http://db.example:4000", "postgresql://sql:sql%20secret@pg.example:6432/my%20db", + Upstream{Address: "pg.example:6432", Host: "pg.example", User: "sql", Password: "sql secret", Database: "my db"}, + }, + { + "PG origin", "postgres://sql:secret@db.example/public", "", + Upstream{Address: "db.example:4003", Host: "db.example", User: "sql", Password: "secret", Database: "public"}, + }, + { + "override wins", "postgres://old:old@old.example/old", "postgres://new:new@new.example/new", + Upstream{Address: "new.example:4003", Host: "new.example", User: "new", Password: "new", Database: "new"}, + }, + } { + t.Run(tt.name, func(t *testing.T) { + o := bo.New() + data, err := yaml.Marshal(map[string]any{"origin_url": tt.origin, "postgres": map[string]string{"upstream_url": tt.override}}) + if err != nil { + t.Fatal(err) + } + if err = yaml.Unmarshal(data, o); err != nil { + t.Fatal(err) + } + c, err := ConfigFromOptions(o, mixedTestEngine{}) + if err != nil { + t.Fatal(err) + } + if c.Upstream != tt.want { + t.Fatalf("got %+v, want %+v", c.Upstream, tt.want) + } + }) + } +} + +func TestMixedUpstreamRejections(t *testing.T) { + for _, tt := range []struct{ name, origin, override string }{ + {"unsupported origin", "mysql://db.example/public", ""}, + {"missing origin host", "http:///v1/sql", ""}, + {"invalid override", "http://db.example:4000", "://bad"}, + {"HTTP override", "http://db.example:4000", "http://db.example:4003"}, + {"missing override host", "http://db.example:4000", "postgres:///public"}, + {"invalid override port", "http://db.example:4000", "postgres://db.example:65536/public"}, + {"malformed credentials", "http://db.example:4000", "postgres://user:secret%zz@db.example/public"}, + } { + t.Run(tt.name, func(t *testing.T) { + o := bo.New() + o.OriginURL = tt.origin + o.Postgres = pgo.New() + o.Postgres.UpstreamURL = tt.override + _, err := ConfigFromOptions(o, mixedTestEngine{}) + if err == nil { + t.Fatal("expected invalid upstream to fail") + } + if strings.Contains(err.Error(), "secret") { + t.Fatal("error exposed upstream credentials") + } + }) + } + o := configTestOptions() + o.OriginURL = "http://http-user:http-secret@db.example:4000/path" + if _, err := ConfigFromOptions(o, mixedTestEngine{}); err == nil { + t.Fatal("HTTP credentials must not enable terminated pgwire authentication") + } + if _, err := ConfigFromOptions(o, testEngine{}); err == nil { + t.Fatal("a native-only engine must still reject HTTP origins") + } +} + +func TestMixedUpstreamReloadAndTLS(t *testing.T) { + o := configTestOptions() + o.OriginURL = "https://http.example:4000/base" + o.Postgres = pgo.New() + o.Postgres.UpstreamURL = "postgres://origin:password@sql.example/public" + o.Postgres.UpstreamTLSMode = pgo.TLSModeVerifyFull + first, err := ConfigFromOptions(o, mixedTestEngine{}) + if err != nil { + t.Fatal(err) + } + if first.Upstream.TLS == nil || first.Upstream.TLS.ServerName != "sql.example" { + t.Fatal("TLS must verify the resolved pgwire host, not the HTTP host") + } + for _, raw := range []string{ + "postgres://origin:rotated@sql.example/public", + "postgres://origin:password@new.example/public", + "postgres://origin:password@sql.example/other", + } { + clone := o.Clone() + clone.Postgres.UpstreamURL = raw + next, err := ConfigFromOptions(clone, mixedTestEngine{}) + if err != nil { + t.Fatal(err) + } + if first.RestartKey == next.RestartKey { + t.Fatal("upstream change did not restart the listener") + } + } + if o.Postgres.UpstreamURL != "postgres://origin:password@sql.example/public" { + t.Fatal("cloning options changed the original pgwire URL") + } +} + +func TestNativeHTTPIsProviderSpecific(t *testing.T) { + a := NewNativeListenerAdapter(NewEngines(testEngine{}, mixedTestEngine{})) + for provider, want := range map[string]bool{"postgres": false, "timescaledb": false, "mixed": true, "unknown": false} { + if got := a.SupportsHTTP(provider); got != want { + t.Errorf("%s: got %t, want %t", provider, got, want) + } + } +} diff --git a/pkg/routing/routing.go b/pkg/routing/routing.go index 65de6e327..16122b7af 100644 --- a/pkg/routing/routing.go +++ b/pkg/routing/routing.go @@ -233,7 +233,7 @@ func RegisterProxyRoutesForListeners(conf *config.Config, clients backends.Backe } routes = append(routes, listenerRoute{r, frontendOptions(conf, name)}) } - if len(o.ListenerNames) == 0 && registry.NativeListeners().GetByProvider(strings.ToLower(o.Provider)) == nil { + if len(o.ListenerNames) == 0 && len(registry.NativeListeners().ForProvider(strings.ToLower(o.Provider))) == 0 { return nil } return routes diff --git a/pkg/routing/routing_test.go b/pkg/routing/routing_test.go index 15a83f74c..e0826e19f 100644 --- a/pkg/routing/routing_test.go +++ b/pkg/routing/routing_test.go @@ -33,6 +33,7 @@ import ( "github.com/trickstercache/trickster/v2/pkg/backends/alb/options" "github.com/trickstercache/trickster/v2/pkg/backends/clickhouse" "github.com/trickstercache/trickster/v2/pkg/backends/graphite" + "github.com/trickstercache/trickster/v2/pkg/backends/greptimedb" "github.com/trickstercache/trickster/v2/pkg/backends/healthcheck" "github.com/trickstercache/trickster/v2/pkg/backends/influxdb" bo "github.com/trickstercache/trickster/v2/pkg/backends/options" @@ -45,6 +46,7 @@ import ( "github.com/trickstercache/trickster/v2/pkg/config" "github.com/trickstercache/trickster/v2/pkg/config/listener" "github.com/trickstercache/trickster/v2/pkg/config/reserved" + configtypes "github.com/trickstercache/trickster/v2/pkg/config/types" "github.com/trickstercache/trickster/v2/pkg/observability/logging" "github.com/trickstercache/trickster/v2/pkg/observability/logging/accesslog" alo "github.com/trickstercache/trickster/v2/pkg/observability/logging/accesslog/options" @@ -850,6 +852,73 @@ func TestBackendRoutesOnMultipleHTTPListeners(t *testing.T) { } } +func TestGreptimeDBRoutesOnlyOnHTTPListeners(t *testing.T) { + var calls atomic.Int64 + origin := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + calls.Add(1) + if user, password, ok := r.BasicAuth(); !ok || user != "grafana" || password != "client-password" { + w.WriteHeader(http.StatusUnauthorized) + return + } + w.Header().Set("X-Origin-URI", r.RequestURI) + w.Header().Set("X-Origin-Method", r.Method) + w.WriteHeader(http.StatusAccepted) + })) + defer origin.Close() + for _, exposeHTTP := range []bool{true, false} { + conf := config.NewConfig() + o := bo.New() + o.Provider, o.OriginURL = providers.GreptimeDB, origin.URL+"/prefix" + o.AuthOptions = &autho.Options{ProxyPreserve: true, Users: configtypes.EnvStringMap{"grafana": "client-password"}} + auth, err := basic.New(map[string]any{"options": o.AuthOptions}) + if err != nil { + t.Fatal(err) + } + o.AuthOptions.Authenticator = auth + o.ListenerNames = []string{"pg"} + if exposeHTTP { + o.ListenerNames = append(o.ListenerNames, "default") + } + if err := o.Initialize("greptime"); err != nil { + t.Fatal(err) + } + conf.Backends = bo.Lookup{o.Name: o} + conf.Listeners["pg"] = &listener.Options{Protocol: listener.ProtocolPostgres, ListenPort: 8489} + caches := registry.LoadCachesFromConfig(conf) + t.Cleanup(func() { registry.CloseCaches(caches) }) + client, err := greptimedb.NewClient(o.Name, o, lm.NewRouter(), caches[o.CacheName], nil, nil) + if err != nil { + t.Fatal(err) + } + o.HTTPClient = client.HTTPClient() + clients := backends.Backends{o.Name: client} + routers := map[string]router.Router{"default": lm.NewRouter(), "pg": lm.NewRouter()} + if err := RegisterProxyRoutesForListeners(conf, clients, routers, nil, caches, nil, false); err != nil { + t.Fatal(err) + } + for _, path := range []string{"/v1/sql?db=public", "/v1/prometheus/api/v1/query_range?query=up", "/v1/influxdb/write", "/v1/loki/api/v1/push"} { + for _, method := range []string{http.MethodGet, http.MethodPost, http.MethodPut, http.MethodDelete, http.MethodPatch, http.MethodOptions, http.MethodHead} { + for _, name := range []string{"default", "pg"} { + for range 2 { + before := calls.Load() + rec := httptest.NewRecorder() + req := httptest.NewRequest(method, "/greptime"+path, nil) + req.SetBasicAuth("grafana", "client-password") + routers[name].ServeHTTP(rec, req) + if name == "default" && exposeHTTP { + if rec.Code != http.StatusAccepted || rec.Header().Get("X-Origin-URI") != "/prefix"+path || rec.Header().Get("X-Origin-Method") != method || calls.Load() != before+1 { + t.Fatalf("%s %s: not relayed exactly once, status=%d headers=%v", method, path, rec.Code, rec.Header()) + } + } else if rec.Code != http.StatusNotFound || calls.Load() != before { + t.Fatalf("HTTP route leaked onto %s (HTTP exposed=%t)", name, exposeHTTP) + } + } + } + } + } + } +} + func TestPassthroughLaneSelection(t *testing.T) { conf := config.NewConfig() o := conf.Backends["default"] diff --git a/pkg/testutil/sqlcompat/corpus.go b/pkg/testutil/sqlcompat/corpus.go new file mode 100644 index 000000000..6b82255ed --- /dev/null +++ b/pkg/testutil/sqlcompat/corpus.go @@ -0,0 +1,277 @@ +/* + * Copyright 2018 The Trickster Authors + * + * 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 sqlcompat runs SQL dialect compatibility corpora in backend tests. +package sqlcompat + +import ( + "encoding/json" + "os" + "slices" + "strings" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/parsing/sqlanalyzer" + "github.com/trickstercache/trickster/v2/pkg/timeseries" +) + +const ( + corpusMinimumInterval = "1m" + corpusPolicyNone = "none" + corpusPolicyRange = "range_independent" + corpusPlaceholder = "TRICKSTER_TS" +) + +type compatibilityCorpus struct { + SchemaVersion int `json:"schema_version"` + CorpusVersion string `json:"corpus_version"` + MinimumInterval string `json:"minimum_interval"` + Omissions []string `json:"omissions"` + Cases []compatibilityCase `json:"cases"` +} + +type compatibilityCase struct { + Name string `json:"name"` + MacroSource []string `json:"macro_source"` + QueryOrigin string `json:"query_origin"` + GrafanaVersion string `json:"grafana_version"` + SessionTimeZone string `json:"session_time_zone"` + ExpandedSQL string `json:"expanded_sql"` + Expected compatibilityExpected `json:"expected"` + Rationale string `json:"rationale"` +} + +type compatibilityExpected struct { + CacheMode string `json:"cache_mode"` + AnalysisReason string `json:"analysis_reason"` + Cadence string `json:"cadence"` + Phase string `json:"phase"` + InputUnit string `json:"input_unit"` + OutputUnit string `json:"output_unit"` + LowerBound string `json:"lower_bound"` + LowerInclusive bool `json:"lower_inclusive"` + UpperBound string `json:"upper_bound"` + UpperInclusive bool `json:"upper_inclusive"` + OpenEnded bool `json:"open_ended"` + OutputColumn string `json:"output_column"` + GroupColumns []string `json:"group_columns"` + CanonicalPolicy string `json:"canonical_policy"` + ExtentRendering bool `json:"extent_rendering"` + CanonicalContains []string `json:"canonical_contains"` + CanonicalExcludes []string `json:"canonical_excludes"` +} + +var ( + corpusModes = map[string]sqlanalyzer.CacheMode{ + "none": sqlanalyzer.CacheModeNone, "object": sqlanalyzer.CacheModeObject, "delta": sqlanalyzer.CacheModeDelta, + } + corpusUnits = map[string]timeseries.FieldDataType{ + "timestamp": timeseries.DateTimeRFC3339Nano, "rfc3339": timeseries.DateTimeRFC3339, + "datetime_sql": timeseries.DateTimeSQL, "unix_seconds": timeseries.DateTimeUnixSecs, + } +) + +func loadCompatibilityCorpus(t testing.TB, path string) compatibilityCorpus { + t.Helper() + data, err := os.ReadFile(path) + if err != nil { + t.Fatal(err) + } + var corpus compatibilityCorpus + if err := json.Unmarshal(data, &corpus); err != nil { + t.Fatal(err) + } + return corpus +} + +// Analyze evaluates a statement in the corpus case's effective session time zone. +type Analyze func(zone, sql string) sqlanalyzer.Analysis + +// Run checks classification, plan facts, and exact extent render/read-back. +func Run(t *testing.T, path string, analyze Analyze) { + t.Helper() + corpus := loadCompatibilityCorpus(t, path) + if corpus.SchemaVersion != 1 || corpus.CorpusVersion == "" || corpus.MinimumInterval != corpusMinimumInterval { + t.Fatalf("invalid corpus header: %+v", corpus) + } + if len(corpus.Omissions) == 0 || len(corpus.Cases) == 0 { + t.Fatal("the corpus must include cases and record what it deliberately leaves out") + } + seen := make(map[string]struct{}, len(corpus.Cases)) + for _, tc := range corpus.Cases { + t.Run(tc.Name, func(t *testing.T) { + if tc.Name == "" || tc.ExpandedSQL == "" || tc.Rationale == "" || tc.SessionTimeZone == "" || + tc.QueryOrigin == "" || tc.Expected.CanonicalPolicy == "" { + t.Fatalf("case lacks required documentation: %+v", tc) + } + if len(tc.MacroSource) > 0 && tc.GrafanaVersion == "" { + t.Fatal("a Grafana macro case must record the Grafana version it was captured from") + } + if _, ok := seen[tc.Name]; ok { + t.Fatalf("duplicate case %q", tc.Name) + } + seen[tc.Name] = struct{}{} + wantMode, ok := corpusModes[tc.Expected.CacheMode] + if !ok { + t.Fatalf("unknown cache mode %q", tc.Expected.CacheMode) + } + analysis := analyze(tc.SessionTimeZone, tc.ExpandedSQL) + if analysis.Mode != wantMode || string(analysis.Reason) != tc.Expected.AnalysisReason { + t.Fatalf("got %s/%s (%v), want %s/%s", analysis.Mode, analysis.Reason, analysis.Err, + tc.Expected.CacheMode, tc.Expected.AnalysisReason) + } + if wantMode != sqlanalyzer.CacheModeDelta { + if analysis.Plan != nil || tc.Expected.ExtentRendering || tc.Expected.CanonicalPolicy != corpusPolicyNone { + t.Fatalf("a case off the delta path must have no renderable plan: %+v", analysis.Plan) + } + return + } + assertCompatibilityPlan(t, tc, analysis.Plan, analyze) + }) + } +} + +func assertCompatibilityPlan(t *testing.T, tc compatibilityCase, plan *sqlanalyzer.QueryPlan, analyze Analyze) { + t.Helper() + want := tc.Expected + if plan == nil || want.CanonicalPolicy != corpusPolicyRange || !want.ExtentRendering { + t.Fatalf("a delta case needs a plan, a range-independent identity and extent rendering: %+v", want) + } + step, err := time.ParseDuration(want.Cadence) + if err != nil { + t.Fatal(err) + } + phase, err := time.ParseDuration(want.Phase) + if err != nil { + t.Fatal(err) + } + lower, err := time.Parse(time.RFC3339Nano, want.LowerBound) + if err != nil { + t.Fatal(err) + } + inputUnit, inputOK := corpusUnits[want.InputUnit] + outputUnit, outputOK := corpusUnits[want.OutputUnit] + if !inputOK || !outputOK { + t.Fatalf("unknown unit in %q / %q", want.InputUnit, want.OutputUnit) + } + if plan.Step != step || plan.Phase != phase || plan.InputUnit != inputUnit || plan.OutputUnit != outputUnit || + plan.OutputColumn != want.OutputColumn || !slices.Equal(plan.GroupColumns, want.GroupColumns) { + t.Fatalf("plan facts: step %v phase %v in %v out %v column %q groups %v", plan.Step, plan.Phase, + plan.InputUnit, plan.OutputUnit, plan.OutputColumn, plan.GroupColumns) + } + if plan.LowerBound == nil || !plan.LowerBound.Value.Equal(lower) || plan.LowerBound.Inclusive != want.LowerInclusive { + t.Fatalf("lower bound %+v, want %s", plan.LowerBound, want.LowerBound) + } + end := lower.Add(step) + if want.OpenEnded { + if plan.UpperBound != nil { + t.Fatalf("expected an open-ended plan, got upper bound %+v", plan.UpperBound) + } + } else { + upper, err := time.Parse(time.RFC3339Nano, want.UpperBound) + if err != nil { + t.Fatal(err) + } + if plan.UpperBound == nil || !plan.UpperBound.Value.Equal(upper) || plan.UpperBound.Inclusive != want.UpperInclusive { + t.Fatalf("upper bound %+v, want %s", plan.UpperBound, want.UpperBound) + } + // every bound lies on the bucket grid, or partial buckets would be cached as whole ones + if !sqlanalyzer.AlignedToBucket(upper, step, phase) { + t.Fatalf("upper bound %s is off the grid", upper) + } + } + if !sqlanalyzer.AlignedToBucket(lower, step, phase) { + t.Fatalf("lower bound %s is off the grid", lower) + } + for _, fragment := range want.CanonicalContains { + if !strings.Contains(plan.CanonicalSQL, fragment) { + t.Errorf("canonical SQL lacks %q: %s", fragment, plan.CanonicalSQL) + } + } + for _, fragment := range want.CanonicalExcludes { + if strings.Contains(plan.CanonicalSQL, fragment) { + t.Errorf("canonical SQL kept %q: %s", fragment, plan.CanonicalSQL) + } + } + rendered, err := plan.RenderExtent(timeseries.Extent{Start: lower, End: end}) + if err != nil { + t.Fatal(err) + } + if strings.Contains(rendered, corpusPlaceholder) || strings.Contains(rendered, "<$") { + t.Fatalf("rendered SQL kept a placeholder: %s", rendered) + } + // what Trickster sends must mean the same statement: same identity, and exactly the extent asked for + again := analyze(tc.SessionTimeZone, rendered) + if again.Mode != sqlanalyzer.CacheModeDelta || again.Plan == nil || again.Plan.CanonicalSQL != plan.CanonicalSQL { + t.Fatalf("rendering changed the statement:\n%s\n%v / %v", rendered, again.Mode, again.Err) + } + if again.Plan.LowerBound == nil || !again.Plan.LowerBound.Value.Equal(lower) || again.Plan.UpperBound == nil || + !again.Plan.UpperBound.Value.Equal(end.Add(step)) { + t.Fatalf("rendered extent reads back as %v..%v, want %v..%v\n%s", again.Plan.LowerBound, + again.Plan.UpperBound, lower, end.Add(step), rendered) + } +} + +// CheckGrafanaMacros requires an explicit case for every bundled PostgreSQL macro. +func CheckGrafanaMacros(t *testing.T, path string) { + t.Helper() + corpus := loadCompatibilityCorpus(t, path) + var sources strings.Builder + for _, tc := range corpus.Cases { + sources.WriteString(strings.Join(tc.MacroSource, "\n")) + sources.WriteByte('\n') + } + for _, macro := range []string{ + "$__time(", "$__timeEpoch(", "$__timeFilter(", "$__timeFrom(", "$__timeTo(", "$__timeGroup(", + "$__timeGroupAlias(", "$__unixEpochFilter(", "$__unixEpochNanoFilter(", "$__unixEpochFrom(", + "$__unixEpochTo(", "$__unixEpochGroup(", "$__unixEpochGroupAlias(", "$__interval", + } { + if !strings.Contains(sources.String(), macro) { + t.Errorf("the corpus does not cover %s", macro) + } + } +} + +// Benchmark measures analysis and immutable concurrent rendering per corpus case. +func Benchmark(b *testing.B, path string, analyze Analyze) { + b.Helper() + corpus := loadCompatibilityCorpus(b, path) + for _, tc := range corpus.Cases { + b.Run("Analyze/"+tc.Expected.CacheMode+"/"+tc.Name, func(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + _ = analyze(tc.SessionTimeZone, tc.ExpandedSQL) + } + }) + plan := analyze(tc.SessionTimeZone, tc.ExpandedSQL).Plan + if plan == nil { + continue + } + extent := timeseries.Extent{Start: plan.LowerBound.Value, End: plan.LowerBound.Value.Add(plan.Step)} + b.Run("Render/"+tc.Name, func(b *testing.B) { + b.ReportAllocs() + b.RunParallel(func(pb *testing.PB) { + for pb.Next() { + if _, err := plan.RenderExtent(extent); err != nil { + b.Error(err) + return + } + } + }) + }) + } +} diff --git a/pkg/timeseries/dataset/arrow/arrow_test.go b/pkg/timeseries/dataset/arrow/arrow_test.go index 40ee4c93e..0dcf834d5 100644 --- a/pkg/timeseries/dataset/arrow/arrow_test.go +++ b/pkg/timeseries/dataset/arrow/arrow_test.go @@ -20,12 +20,12 @@ import ( "errors" "fmt" "math" - "math/rand" "reflect" "testing" "github.com/trickstercache/trickster/v2/pkg/timeseries" "github.com/trickstercache/trickster/v2/pkg/timeseries/dataset" + "github.com/trickstercache/trickster/v2/pkg/util/weak/weaktest" "github.com/apache/arrow-go/v18/arrow" "github.com/apache/arrow-go/v18/arrow/array" @@ -369,7 +369,7 @@ func TestChunking(t *testing.T) { // TestRoundTripRandomized property-tests the round trip over randomized // schemas and values drawn from the supported type pool. func TestRoundTripRandomized(t *testing.T) { - rng := rand.New(rand.NewSource(42)) + rng := weaktest.NewRand(42, 0) valueTypes := []arrow.DataType{ arrow.FixedWidthTypes.Boolean, arrow.PrimitiveTypes.Int8, arrow.PrimitiveTypes.Int16, @@ -383,7 +383,7 @@ func TestRoundTripRandomized(t *testing.T) { randomCell := func(dt arrow.DataType) any { switch dt.ID() { case arrow.BOOL: - return rng.Intn(2) == 0 + return rng.IntN(2) == 0 case arrow.INT8: return int64(int8(rng.Int())) case arrow.INT16: @@ -391,7 +391,7 @@ func TestRoundTripRandomized(t *testing.T) { case arrow.INT32: return int64(int32(rng.Int())) case arrow.INT64, arrow.TIMESTAMP: - return rng.Int63() - rng.Int63() + return rng.Int64() - rng.Int64() case arrow.UINT8: return int64(uint8(rng.Int())) case arrow.UINT16: @@ -405,7 +405,7 @@ func TestRoundTripRandomized(t *testing.T) { case arrow.FLOAT64: return rng.NormFloat64() default: // string-ish - return fmt.Sprintf("s%d", rng.Intn(1000)) + return fmt.Sprintf("s%d", rng.IntN(1000)) } } @@ -414,17 +414,17 @@ func TestRoundTripRandomized(t *testing.T) { {Name: "time", Type: &arrow.TimestampType{Unit: arrow.Nanosecond, TimeZone: "UTC"}}, {Name: "tag", Type: arrow.BinaryTypes.String}, } - for i := range 1 + rng.Intn(6) { + for i := range 1 + rng.IntN(6) { fields = append(fields, arrow.Field{ Name: fmt.Sprintf("v%d", i), - Type: valueTypes[rng.Intn(len(valueTypes))], + Type: valueTypes[rng.IntN(len(valueTypes))], Nullable: true, }) } schema := arrow.NewSchema(fields, nil) - tagValues := []string{"a", "b", "c"}[:1+rng.Intn(3)] - rowCount := rng.Intn(50) + tagValues := []string{"a", "b", "c"}[:1+rng.IntN(3)] + rowCount := rng.IntN(50) rows := make([][]any, rowCount) for r := range rows { cells := make([]any, len(fields)) diff --git a/pkg/timeseries/dataset/builder.go b/pkg/timeseries/dataset/builder.go new file mode 100644 index 000000000..b97a9d8b8 --- /dev/null +++ b/pkg/timeseries/dataset/builder.go @@ -0,0 +1,468 @@ +/* + * Copyright 2018 The Trickster Authors + * + * 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 dataset + +import ( + "errors" + "fmt" + "slices" + "strconv" + + "github.com/trickstercache/trickster/v2/pkg/timeseries" + "github.com/trickstercache/trickster/v2/pkg/timeseries/epoch" +) + +// DuplicatePolicy controls how a Builder treats points in one series that share an epoch. +type DuplicatePolicy byte + +const ( + // DuplicatesKeep retains every point, including those sharing an epoch. + DuplicatesKeep DuplicatePolicy = iota + // DuplicatesFirstWins retains only the first-received point for each epoch. + DuplicatesFirstWins + // DuplicatesLastWins retains only the last-received point for each epoch. + DuplicatesLastWins + // DuplicatesError fails the build when a series receives two points for one epoch. + DuplicatesError +) + +var ( + // ErrDuplicateEpoch indicates a series received more than one point for an epoch. + ErrDuplicateEpoch = fmt.Errorf("%w: duplicate epoch in series", timeseries.ErrInvalidBody) + // ErrInvalidRow indicates a row or point that does not fit the Builder's fields. + ErrInvalidRow = fmt.Errorf("%w: row does not match the dataset fields", timeseries.ErrInvalidBody) + // ErrBuilderFinished indicates a Builder was used after Finish. + ErrBuilderFinished = errors.New("dataset builder already finished") +) + +// BuilderOptions configures a Builder. +type BuilderOptions struct { + // Fields describes the timestamp, tag and value fields of each row in row mode. + Fields timeseries.SeriesFields + // SeriesName is the header name of each series created in row mode. + SeriesName string + // QueryStatement is the header query statement of each series created in row mode. + QueryStatement string + // Duplicates controls how points sharing an epoch within a series are handled. + Duplicates DuplicatePolicy + // SortSeries sorts each result's series by their tags when the build finishes. + SortSeries bool + // TagString converts a raw tag value, valid only during the call, to its Tags entry; + // rows whose converted tags match share a series. When nil, raw bytes are used as-is. + TagString func(fd timeseries.FieldDefinition, raw []byte) string +} + +// Builder assembles a DataSet in a single pass from rows or points in any order. +// A Builder is not safe for concurrent use. +type Builder struct { + trq *timeseries.TimeRangeQuery + opts BuilderOptions + results []*resultBuild + result *resultBuild + current *seriesBuild + row RowBuilder + arena []any + chunk int + finished bool +} + +// RowBuilder accumulates one row or point for a Builder; it is reused by each +// call to Builder.Row and is only valid until the next call. +type RowBuilder struct { + b *Builder + epoch epoch.Epoch + hasEpoch bool + invalid bool + values []any + tagBuf []byte + tagOff []int + key []byte +} + +type resultBuild struct { + r *Result + series []*seriesBuild + lookup map[string]*seriesBuild + index *seriesIndex[*seriesBuild] +} + +type seriesBuild struct { + s *Series + expected int + unordered bool +} + +const ( + minValueChunk = 64 + maxValueChunk = 4096 + pointOverhead = 40 // a Point's Epoch, Size and Values slice header + valueOverhead = 16 // the interface header of each value +) + +// NewBuilder returns a Builder for the provided query and options. +func NewBuilder(trq *timeseries.TimeRangeQuery, opts BuilderOptions) *Builder { + b := &Builder{trq: trq, opts: opts} + b.row.b = b + b.row.tagOff = make([]int, 2*len(opts.Fields.Tags)) + b.row.reset() + return b +} + +// SetResult ends any open series and directs subsequent rows and series to the +// result with the provided statement ID and name, creating it if needed. +func (b *Builder) SetResult(statementID int, name string) { + if b.finished { + return + } + b.current = nil + for _, rb := range b.results { + if rb.r.StatementID == statementID && rb.r.Name == name { + b.result = rb + return + } + } + b.result = &resultBuild{ + r: &Result{StatementID: statementID, Name: name, SeriesList: SeriesList{}}, + lookup: make(map[string]*seriesBuild), + index: newSeriesIndex(0, builtHeader), + } + b.results = append(b.results, b.result) +} + +// StartSeries switches the Builder to series mode: rows and points go to the series +// with an identical header, created if needed, until EndSeries or the next StartSeries. +func (b *Builder) StartSeries(h SeriesHeader) { + if b.finished { + return + } + b.current = b.seriesFor(b.currentResult(), h, false) +} + +// EndSeries returns the Builder to row mode, where rows are grouped into series by their tags. +func (b *Builder) EndSeries() { + b.current = nil +} + +// Row returns the Builder's reusable RowBuilder, reset for a new row. +func (b *Builder) Row() *RowBuilder { + b.row.reset() + return &b.row +} + +// AppendPoint appends p to the series opened by StartSeries. When p.Size is 0, +// it is set by PointSize. +func (b *Builder) AppendPoint(p Point) error { + if b.finished { + return ErrBuilderFinished + } + sb := b.current + if sb == nil || (sb.expected > 0 && len(p.Values) != sb.expected) { + return ErrInvalidRow + } + if p.Size == 0 { + p.Size = PointSize(p.Values) + } + return b.appendPoint(sb, p.Epoch, p.Values, p.Size, false) +} + +// Finish sorts only the series that arrived out of order, applies the duplicate +// policy, and returns the DataSet. The Builder cannot be used afterward. +func (b *Builder) Finish() (*DataSet, error) { + if b.finished { + return nil, ErrBuilderFinished + } + b.currentResult() + b.finished = true + b.current = nil + ds := &DataSet{TimeRangeQuery: b.trq, Results: make(Results, len(b.results))} + if b.trq != nil { + ds.ExtentList = timeseries.ExtentList{b.trq.Extent} + } + for i, rb := range b.results { + rb.r.SeriesList = make(SeriesList, len(rb.series)) + for j, sb := range rb.series { + if err := b.finishSeries(sb); err != nil { + return nil, err + } + rb.r.SeriesList[j] = sb.s + } + if b.opts.SortSeries { + rb.r.SeriesList.SortByTags() + } + ds.Results[i] = rb.r + } + return ds, nil +} + +// PointSize returns the estimated memory, in bytes, of a Point holding values. +func PointSize(values []any) int { + n := pointOverhead + for _, v := range values { + n += valueSize(v) + } + return n +} + +func valueSize(v any) int { + switch t := v.(type) { + case nil: + return valueOverhead + case string: + return valueOverhead + len(t) + case []byte: + return valueOverhead + len(t) + case bool, int8, uint8: + return valueOverhead + 1 + case int16, uint16: + return valueOverhead + 2 + case int32, uint32, float32: + return valueOverhead + 4 + } + return valueOverhead + 8 +} + +// SetEpoch sets the row's timestamp. +func (r *RowBuilder) SetEpoch(e epoch.Epoch) { + r.epoch = e + r.hasEpoch = true +} + +// SetTag sets the raw value of tag field i, per BuilderOptions.Fields.Tags; it +// is copied. Unset tags are omitted from the series' Tags. +func (r *RowBuilder) SetTag(i int, raw []byte) { + if i < 0 || 2*i >= len(r.tagOff) { + r.invalid = true + return + } + r.tagOff[2*i] = len(r.tagBuf) + r.tagBuf = append(r.tagBuf, raw...) + r.tagOff[2*i+1] = len(r.tagBuf) +} + +// AddValue appends the next value, per BuilderOptions.Fields.Values or the +// open series' ValueFieldsList. +func (r *RowBuilder) AddValue(v any) { + r.values = append(r.values, v) +} + +// Commit adds the row to its series: the open series in series mode, or else the +// series matching the row's tags, which is created on first use. +func (r *RowBuilder) Commit() error { + b := r.b + if b.finished { + return ErrBuilderFinished + } + if !r.hasEpoch || r.invalid { + return ErrInvalidRow + } + sb := b.current + expected := len(b.opts.Fields.Values) + if sb != nil { + if slices.ContainsFunc(r.tagOff, func(o int) bool { return o >= 0 }) { + return ErrInvalidRow + } + expected = sb.expected + } + if expected > 0 && len(r.values) != expected { + return ErrInvalidRow + } + if sb == nil { + sb = r.series() + } + var values []any + if len(r.values) > 0 { + values = b.allocValues(len(r.values)) + copy(values, r.values) + } + err := b.appendPoint(sb, r.epoch, values, PointSize(values), true) + r.reset() + return err +} + +func (r *RowBuilder) reset() { + r.hasEpoch = false + r.invalid = false + clear(r.values) + r.values = r.values[:0] + r.tagBuf = r.tagBuf[:0] + for i := range r.tagOff { + r.tagOff[i] = -1 + } +} + +func (r *RowBuilder) series() *seriesBuild { + r.key = r.key[:0] + for i := 0; i < len(r.tagOff); i += 2 { + start, end := r.tagOff[i], r.tagOff[i+1] + if start < 0 { + r.key = append(r.key, '-') + continue + } + r.key = strconv.AppendInt(r.key, int64(end-start), 10) + r.key = append(r.key, ':') + r.key = append(r.key, r.tagBuf[start:end]...) + } + rb := r.b.currentResult() + if sb, ok := rb.lookup[string(r.key)]; ok { + return sb + } + opts := &r.b.opts + tags := make(Tags, len(opts.Fields.Tags)) + for i, fd := range opts.Fields.Tags { + start, end := r.tagOff[2*i], r.tagOff[2*i+1] + if start < 0 { + continue + } + raw := r.tagBuf[start:end] + if opts.TagString != nil { + tags[fd.Name] = opts.TagString(fd, raw) + } else { + tags[fd.Name] = string(raw) + } + } + // a new raw encoding can still name an existing series, as "a" and "\u0061" do in JSON + sb := r.b.seriesFor(rb, SeriesHeader{ + Name: opts.SeriesName, + Tags: tags, + TimestampField: opts.Fields.Timestamp, + TagFieldsList: opts.Fields.Tags, + ValueFieldsList: opts.Fields.Values, + UntrackedFieldsList: opts.Fields.Untracked, + QueryStatement: opts.QueryStatement, + }, true) + rb.lookup[string(r.key)] = sb + return sb +} + +func (b *Builder) seriesFor(rb *resultBuild, h SeriesHeader, cloneFields bool) *seriesBuild { + // series are matched the way merges match them, so no two can later merge as one + hash := h.CalculateHashWithQueryStatement(h.QueryStatement) + if sb, ok := rb.index.find(hash, &h); ok { + return sb + } + if cloneFields { + h.TagFieldsList = slices.Clone(h.TagFieldsList) + h.ValueFieldsList = slices.Clone(h.ValueFieldsList) + h.UntrackedFieldsList = slices.Clone(h.UntrackedFieldsList) + } + sb := &seriesBuild{s: &Series{Header: h}, expected: len(h.ValueFieldsList)} + rb.index.add(hash, sb) + rb.series = append(rb.series, sb) + return sb +} + +func builtHeader(sb *seriesBuild) *SeriesHeader { + return &sb.s.Header +} + +func (b *Builder) currentResult() *resultBuild { + if b.result == nil { + b.SetResult(0, "") + } + return b.result +} + +func (b *Builder) allocValues(n int) []any { + // points share chunked backing arrays to avoid an allocation per point + if cap(b.arena)-len(b.arena) < n { + b.chunk = min(max(2*b.chunk, minValueChunk), maxValueChunk) + b.arena = make([]any, 0, max(b.chunk, n)) + } + l := len(b.arena) + b.arena = b.arena[:l+n] + return b.arena[l : l+n : l+n] +} + +func (b *Builder) appendPoint(sb *seriesBuild, e epoch.Epoch, values []any, size int, owned bool) error { + s := sb.s + if n := len(s.Points); n > 0 && !sb.unordered { + last := &s.Points[n-1] + switch { + case e < last.Epoch: + sb.unordered = true + case e == last.Epoch: + switch b.opts.Duplicates { + case DuplicatesError: + return ErrDuplicateEpoch + case DuplicatesFirstWins: + if owned { + b.releaseValues(values) + } + return nil + case DuplicatesLastWins: + s.PointSize += int64(size - last.Size) + last.Values, last.Size = values, size + return nil + } + } + } + s.Points = append(s.Points, Point{Epoch: e, Size: size, Values: values}) + s.PointSize += int64(size) + return nil +} + +func (b *Builder) releaseValues(values []any) { + // only the most recent allocation can be returned to its chunk + n := len(values) + l := len(b.arena) + if n == 0 || l < n || &b.arena[l-n] != &values[0] { + return + } + clear(values) + b.arena = b.arena[:l-n] +} + +func (b *Builder) finishSeries(sb *seriesBuild) error { + s := sb.s + if sb.unordered { + // a stable sort keeps arrival order among equal epochs for the duplicate policy + slices.SortStableFunc(s.Points, pointCmp) + if b.opts.Duplicates != DuplicatesKeep { + if err := b.dedupe(s); err != nil { + return err + } + } + } + s.Header.CalculateSize() + return nil +} + +func (b *Builder) dedupe(s *Series) error { + pts := s.Points + k := 0 + for i := 1; i < len(pts); i++ { + if pts[i].Epoch != pts[k].Epoch { + k++ + pts[k] = pts[i] + continue + } + switch b.opts.Duplicates { + case DuplicatesError: + return ErrDuplicateEpoch + case DuplicatesLastWins: + s.PointSize -= int64(pts[k].Size) + pts[k] = pts[i] + default: + s.PointSize -= int64(pts[i].Size) + } + } + if len(pts) > 0 { + clear(pts[k+1:]) + s.Points = pts[:k+1] + } + return nil +} diff --git a/pkg/timeseries/dataset/builder_test.go b/pkg/timeseries/dataset/builder_test.go new file mode 100644 index 000000000..5a8d0e5e5 --- /dev/null +++ b/pkg/timeseries/dataset/builder_test.go @@ -0,0 +1,584 @@ +/* + * Copyright 2018 The Trickster Authors + * + * 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 dataset + +import ( + "strings" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/timeseries" + "github.com/trickstercache/trickster/v2/pkg/timeseries/epoch" + + "github.com/stretchr/testify/require" +) + +func testBuilderFields() timeseries.SeriesFields { + return timeseries.SeriesFields{ + Timestamp: timeseries.FieldDefinition{Name: "t", DataType: timeseries.DateTimeUnixMilli, Role: timeseries.RoleTimestamp}, + Tags: timeseries.FieldDefinitions{ + {Name: "host", DataType: timeseries.String, Role: timeseries.RoleTag, OutputPosition: 1}, + {Name: "dc", DataType: timeseries.String, Role: timeseries.RoleTag, OutputPosition: 2}, + }, + Values: timeseries.FieldDefinitions{ + {Name: "v", DataType: timeseries.Float64, Role: timeseries.RoleValue, OutputPosition: 3}, + }, + } +} + +func testBuilderTRQ() *timeseries.TimeRangeQuery { + return ×eries.TimeRangeQuery{ + Statement: "select", + Extent: timeseries.Extent{Start: time.Unix(0, 0), End: time.Unix(100, 0)}, + Step: time.Second, + } +} + +type testRow struct { + e epoch.Epoch + host string + dc string + v any +} + +func commitRows(t *testing.T, b *Builder, rows ...testRow) { + t.Helper() + for _, tr := range rows { + r := b.Row() + r.SetEpoch(tr.e) + r.SetTag(1, []byte(tr.dc)) + r.SetTag(0, []byte(tr.host)) + r.AddValue(tr.v) + require.NoError(t, r.Commit()) + } +} + +func pointEpochs(s *Series) []epoch.Epoch { + out := make([]epoch.Epoch, len(s.Points)) + for i, p := range s.Points { + out[i] = p.Epoch + } + return out +} + +func pointValues(s *Series) []any { + out := make([]any, len(s.Points)) + for i, p := range s.Points { + out[i] = p.Values[0] + } + return out +} + +func requireSizes(t *testing.T, s *Series) { + t.Helper() + var total int64 + for _, p := range s.Points { + require.Equal(t, PointSize(p.Values), p.Size) + total += int64(p.Size) + } + require.Equal(t, total, s.PointSize) + require.Positive(t, s.Header.Size) +} + +func TestBuilderRowMode(t *testing.T) { + trq := testBuilderTRQ() + fields := testBuilderFields() + b := NewBuilder(trq, BuilderOptions{Fields: fields, SeriesName: "sql", QueryStatement: trq.Statement}) + commitRows(t, b, + testRow{e: 1, host: "a", dc: "x", v: 1.0}, + testRow{e: 1, host: "b", dc: "x", v: 2.0}, + testRow{e: 2, host: "a", dc: "x", v: 3.0}, + testRow{e: 3, host: "a", dc: "x", v: 4.0}, + testRow{e: 2, host: "b", dc: "x", v: 5.0}, + ) + ds, err := b.Finish() + require.NoError(t, err) + require.Same(t, trq, ds.TimeRangeQuery) + require.Equal(t, timeseries.ExtentList{trq.Extent}, ds.ExtentList) + require.Len(t, ds.Results, 1) + sl := ds.Results[0].SeriesList + require.Len(t, sl, 2) + require.Equal(t, Tags{"host": "a", "dc": "x"}, sl[0].Header.Tags) + require.Equal(t, Tags{"host": "b", "dc": "x"}, sl[1].Header.Tags) + require.Equal(t, []epoch.Epoch{1, 2, 3}, pointEpochs(sl[0])) + require.Equal(t, []any{1.0, 3.0, 4.0}, pointValues(sl[0])) + require.Equal(t, []epoch.Epoch{1, 2}, pointEpochs(sl[1])) + for _, s := range sl { + require.Equal(t, "sql", s.Header.Name) + require.Equal(t, "select", s.Header.QueryStatement) + require.Equal(t, fields.Timestamp, s.Header.TimestampField) + require.Equal(t, fields.Tags, s.Header.TagFieldsList) + require.Equal(t, fields.Values, s.Header.ValueFieldsList) + requireSizes(t, s) + } + // each series owns its field definitions + sl[0].Header.ValueFieldsList[0].DataType = timeseries.Int64 + require.Equal(t, timeseries.Float64, sl[1].Header.ValueFieldsList[0].DataType) + require.Equal(t, timeseries.Float64, fields.Values[0].DataType) +} + +func TestBuilderSortsOnlyUnorderedSeries(t *testing.T) { + b := NewBuilder(nil, BuilderOptions{Fields: testBuilderFields()}) + commitRows(t, b, + testRow{e: 3, host: "a", v: 1.0}, + testRow{e: 1, host: "b", v: 2.0}, + testRow{e: 1, host: "a", v: 3.0}, + testRow{e: 2, host: "b", v: 4.0}, + testRow{e: 2, host: "a", v: 5.0}, + ) + require.True(t, b.results[0].series[0].unordered) + require.False(t, b.results[0].series[1].unordered) + ds, err := b.Finish() + require.NoError(t, err) + require.Nil(t, ds.ExtentList) + sl := ds.Results[0].SeriesList + require.Equal(t, []epoch.Epoch{1, 2, 3}, pointEpochs(sl[0])) + require.Equal(t, []any{3.0, 5.0, 1.0}, pointValues(sl[0])) + require.Equal(t, []epoch.Epoch{1, 2}, pointEpochs(sl[1])) +} + +func TestBuilderDuplicatePolicies(t *testing.T) { + ordered := []testRow{ + {e: 1, host: "a", v: "a1"}, {e: 2, host: "a", v: "a2"}, {e: 2, host: "a", v: "a2b"}, + {e: 2, host: "a", v: "a2c"}, {e: 3, host: "a", v: "a3"}, + } + unordered := []testRow{ + {e: 2, host: "a", v: "a2"}, {e: 1, host: "a", v: "a1"}, {e: 2, host: "a", v: "a2b"}, + {e: 3, host: "a", v: "a3"}, {e: 2, host: "a", v: "a2c"}, + } + tests := []struct { + name string + policy DuplicatePolicy + rows []testRow + want []any + }{ + {"keep ordered", DuplicatesKeep, ordered, []any{"a1", "a2", "a2b", "a2c", "a3"}}, + {"keep unordered", DuplicatesKeep, unordered, []any{"a1", "a2", "a2b", "a2c", "a3"}}, + {"first ordered", DuplicatesFirstWins, ordered, []any{"a1", "a2", "a3"}}, + {"first unordered", DuplicatesFirstWins, unordered, []any{"a1", "a2", "a3"}}, + {"last ordered", DuplicatesLastWins, ordered, []any{"a1", "a2c", "a3"}}, + {"last unordered", DuplicatesLastWins, unordered, []any{"a1", "a2c", "a3"}}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + b := NewBuilder(nil, BuilderOptions{Fields: testBuilderFields(), Duplicates: test.policy}) + commitRows(t, b, test.rows...) + ds, err := b.Finish() + require.NoError(t, err) + s := ds.Results[0].SeriesList[0] + require.Equal(t, test.want, pointValues(s)) + requireSizes(t, s) + }) + } +} + +func TestBuilderDuplicateError(t *testing.T) { + b := NewBuilder(nil, BuilderOptions{Fields: testBuilderFields(), Duplicates: DuplicatesError}) + commitRows(t, b, testRow{e: 1, host: "a", v: 1.0}) + r := b.Row() + r.SetEpoch(1) + r.SetTag(0, []byte("a")) + r.SetTag(1, nil) + r.AddValue(2.0) + require.ErrorIs(t, r.Commit(), ErrDuplicateEpoch) + require.ErrorIs(t, ErrDuplicateEpoch, timeseries.ErrInvalidBody) + + b = NewBuilder(nil, BuilderOptions{Fields: testBuilderFields(), Duplicates: DuplicatesError}) + commitRows(t, b, + testRow{e: 2, host: "a", v: 1.0}, + testRow{e: 1, host: "a", v: 2.0}, + testRow{e: 2, host: "a", v: 3.0}, + ) + _, err := b.Finish() + require.ErrorIs(t, err, ErrDuplicateEpoch) +} + +func TestBuilderFirstWinsReleasesValues(t *testing.T) { + b := NewBuilder(nil, BuilderOptions{Fields: testBuilderFields(), Duplicates: DuplicatesFirstWins}) + commitRows(t, b, testRow{e: 1, host: "a", v: 1.0}) + used := len(b.arena) + commitRows(t, b, testRow{e: 1, host: "a", v: 2.0}) + require.Len(t, b.arena, used) + // values that are not the latest allocation cannot be released + b.releaseValues([]any{1}) + b.releaseValues(nil) + require.Len(t, b.arena, used) +} + +func TestBuilderTags(t *testing.T) { + var calls int + b := NewBuilder(nil, BuilderOptions{ + Fields: testBuilderFields(), + TagString: func(fd timeseries.FieldDefinition, raw []byte) string { + calls++ + return fd.Name + "=" + strings.ToUpper(string(raw)) + }, + SortSeries: true, + }) + add := func(host []byte, setDC bool) { + r := b.Row() + r.SetEpoch(1) + if host != nil { + r.SetTag(0, host) + } + if setDC { + r.SetTag(1, []byte("x")) + } + r.AddValue(1.0) + require.NoError(t, r.Commit()) + } + add([]byte("b"), true) + add([]byte("b"), true) + add(nil, true) + add([]byte(""), true) + add([]byte("a"), false) + // a tag set twice keeps its last value + r := b.Row() + r.SetEpoch(2) + r.SetTag(0, []byte("zzz")) + r.SetTag(0, []byte("b")) + r.SetTag(1, []byte("x")) + r.AddValue(1.0) + require.NoError(t, r.Commit()) + + ds, err := b.Finish() + require.NoError(t, err) + require.Equal(t, 6, calls) + sl := ds.Results[0].SeriesList + require.Len(t, sl, 4) + got := make([]Tags, len(sl)) + for i, s := range sl { + got[i] = s.Header.Tags + } + // SortByTags orders series by their tags' JSON encoding + require.Equal(t, []Tags{ + {"dc": "dc=X", "host": "host="}, + {"dc": "dc=X", "host": "host=B"}, + {"dc": "dc=X"}, + {"host": "host=A"}, + }, got) + require.Equal(t, []epoch.Epoch{1, 1, 2}, pointEpochs(sl[1])) +} + +func TestBuilderInvalidRows(t *testing.T) { + b := NewBuilder(nil, BuilderOptions{Fields: testBuilderFields()}) + r := b.Row() + r.SetTag(0, []byte("a")) + r.AddValue(1.0) + require.ErrorIs(t, r.Commit(), ErrInvalidRow) + + r = b.Row() + r.SetEpoch(1) + r.SetTag(2, []byte("a")) + r.AddValue(1.0) + require.ErrorIs(t, r.Commit(), ErrInvalidRow) + + r = b.Row() + r.SetEpoch(1) + r.SetTag(-1, []byte("a")) + r.AddValue(1.0) + require.ErrorIs(t, r.Commit(), ErrInvalidRow) + + r = b.Row() + r.SetEpoch(1) + r.SetTag(0, []byte("a")) + r.AddValue(1.0) + r.AddValue(2.0) + require.ErrorIs(t, r.Commit(), ErrInvalidRow) + require.ErrorIs(t, ErrInvalidRow, timeseries.ErrInvalidBody) + + ds, err := b.Finish() + require.NoError(t, err) + require.Empty(t, ds.Results[0].SeriesList) +} + +func TestBuilderSeriesMode(t *testing.T) { + b := NewBuilder(nil, BuilderOptions{}) + h := SeriesHeader{ + Name: "up", + Tags: Tags{"job": "a"}, + ValueFieldsList: timeseries.FieldDefinitions{{Name: "value", DataType: timeseries.String}}, + } + b.StartSeries(h) + for _, e := range []epoch.Epoch{2, 1, 3} { + r := b.Row() + r.SetEpoch(e) + r.AddValue("1") + require.NoError(t, r.Commit()) + } + require.NoError(t, b.AppendPoint(Point{Epoch: 4, Values: []any{"42"}})) + require.NoError(t, b.AppendPoint(Point{Epoch: 5, Values: []any{"7"}, Size: 99})) + + r := b.Row() + r.SetEpoch(6) + r.AddValue("1") + r.AddValue("2") + require.ErrorIs(t, r.Commit(), ErrInvalidRow) + require.ErrorIs(t, b.AppendPoint(Point{Epoch: 6}), ErrInvalidRow) + + // tags cannot be set on rows committed to an open series + b2 := NewBuilder(nil, BuilderOptions{Fields: testBuilderFields()}) + b2.StartSeries(SeriesHeader{Name: "x"}) + r2 := b2.Row() + r2.SetEpoch(1) + r2.SetTag(0, []byte("a")) + require.ErrorIs(t, r2.Commit(), ErrInvalidRow) + // a series without value fields accepts any number of values + r2 = b2.Row() + r2.SetEpoch(1) + require.NoError(t, r2.Commit()) + + b.StartSeries(SeriesHeader{Name: "down", Tags: Tags{"job": "b"}}) + require.NoError(t, b.AppendPoint(Point{Epoch: 1, Values: []any{"1", "2"}})) + b.EndSeries() + require.ErrorIs(t, b.AppendPoint(Point{Epoch: 1}), ErrInvalidRow) + + ds, err := b.Finish() + require.NoError(t, err) + sl := ds.Results[0].SeriesList + require.Len(t, sl, 2) + require.Equal(t, "up", sl[0].Header.Name) + require.Equal(t, []epoch.Epoch{1, 2, 3, 4, 5}, pointEpochs(sl[0])) + require.Equal(t, PointSize([]any{"42"}), sl[0].Points[3].Size) + require.Equal(t, 99, sl[0].Points[4].Size) + require.Equal(t, "down", sl[1].Header.Name) +} + +func TestBuilderSeriesModeDuplicates(t *testing.T) { + b := NewBuilder(nil, BuilderOptions{Duplicates: DuplicatesLastWins}) + b.StartSeries(SeriesHeader{Name: "s"}) + require.NoError(t, b.AppendPoint(Point{Epoch: 1, Values: []any{"a"}})) + require.NoError(t, b.AppendPoint(Point{Epoch: 1, Values: []any{"bb"}})) + b.StartSeries(SeriesHeader{Name: "t"}) + require.NoError(t, b.AppendPoint(Point{Epoch: 1, Values: []any{"a"}})) + ds, err := b.Finish() + require.NoError(t, err) + s := ds.Results[0].SeriesList[0] + require.Equal(t, []any{"bb"}, pointValues(s)) + requireSizes(t, s) + + b = NewBuilder(nil, BuilderOptions{Duplicates: DuplicatesFirstWins}) + b.StartSeries(SeriesHeader{Name: "s"}) + require.NoError(t, b.AppendPoint(Point{Epoch: 1, Values: []any{"a"}})) + require.NoError(t, b.AppendPoint(Point{Epoch: 1, Values: []any{"bb"}})) + r := b.Row() + r.SetEpoch(1) + require.NoError(t, r.Commit()) + ds, err = b.Finish() + require.NoError(t, err) + require.Equal(t, []any{"a"}, pointValues(ds.Results[0].SeriesList[0])) +} + +func TestBuilderResults(t *testing.T) { + b := NewBuilder(nil, BuilderOptions{Fields: testBuilderFields()}) + b.SetResult(1, "first") + commitRows(t, b, testRow{e: 1, host: "a", v: 1.0}) + b.StartSeries(SeriesHeader{Name: "open"}) + b.SetResult(2, "second") + require.ErrorIs(t, b.AppendPoint(Point{Epoch: 1}), ErrInvalidRow) + commitRows(t, b, testRow{e: 1, host: "a", v: 2.0}) + b.SetResult(1, "first") + commitRows(t, b, testRow{e: 2, host: "a", v: 3.0}) + b.SetResult(3, "empty") + + ds, err := b.Finish() + require.NoError(t, err) + require.Len(t, ds.Results, 3) + r1, r2, r3 := ds.Results[0], ds.Results[1], ds.Results[2] + require.Equal(t, 1, r1.StatementID) + require.Equal(t, "first", r1.Name) + require.Len(t, r1.SeriesList, 2) + require.Equal(t, "open", r1.SeriesList[1].Header.Name) + require.Equal(t, []any{1.0, 3.0}, pointValues(r1.SeriesList[0])) + require.Equal(t, 2, r2.StatementID) + require.Equal(t, []any{2.0}, pointValues(r2.SeriesList[0])) + require.Equal(t, "empty", r3.Name) + require.NotNil(t, r3.SeriesList) + require.Empty(t, r3.SeriesList) +} + +func TestBuilderEmpty(t *testing.T) { + ds, err := NewBuilder(testBuilderTRQ(), BuilderOptions{}).Finish() + require.NoError(t, err) + require.Len(t, ds.Results, 1) + require.NotNil(t, ds.Results[0].SeriesList) + require.Empty(t, ds.Results[0].SeriesList) + require.Zero(t, ds.ValueCount()) +} + +func TestBuilderFinished(t *testing.T) { + b := NewBuilder(nil, BuilderOptions{Fields: testBuilderFields()}) + _, err := b.Finish() + require.NoError(t, err) + _, err = b.Finish() + require.ErrorIs(t, err, ErrBuilderFinished) + b.SetResult(9, "late") + b.StartSeries(SeriesHeader{Name: "late"}) + require.Len(t, b.results, 1) + require.ErrorIs(t, b.AppendPoint(Point{}), ErrBuilderFinished) + r := b.Row() + r.SetEpoch(1) + require.ErrorIs(t, r.Commit(), ErrBuilderFinished) +} + +func TestBuilderValueChunks(t *testing.T) { + b := NewBuilder(nil, BuilderOptions{}) + b.StartSeries(SeriesHeader{Name: "s"}) + for i := range 3 * minValueChunk { + r := b.Row() + r.SetEpoch(epoch.Epoch(i)) + r.AddValue(int64(i)) + require.NoError(t, r.Commit()) + } + require.Equal(t, 2*minValueChunk, b.chunk) + wide := b.Row() + wide.SetEpoch(epoch.Epoch(3 * minValueChunk)) + for i := range 2 * maxValueChunk { + wide.AddValue(i) + } + require.NoError(t, wide.Commit()) + ds, err := b.Finish() + require.NoError(t, err) + pts := ds.Results[0].SeriesList[0].Points + require.Len(t, pts, 3*minValueChunk+1) + for i, p := range pts[:3*minValueChunk] { + require.Equal(t, []any{int64(i)}, p.Values) + require.Equal(t, 1, cap(p.Values)) + } + require.Len(t, pts[3*minValueChunk].Values, 2*maxValueChunk) +} + +func TestPointSize(t *testing.T) { + require.Equal(t, pointOverhead, PointSize(nil)) + values := []any{nil, "abc", []byte("ab"), true, int8(1), uint8(1), int16(1), uint16(1), + int32(1), uint32(1), float32(1), int64(1), 1.0, uint64(1), 1} + want := pointOverhead + len(values)*valueOverhead + 3 + 2 + 3 + 4 + 12 + 32 + require.Equal(t, want, PointSize(values)) +} + +func BenchmarkBuilderRows(b *testing.B) { + hosts := [][]byte{[]byte("host-a"), []byte("host-b"), []byte("host-c"), []byte("host-d")} + dc := []byte("dc-1") + fields := testBuilderFields() + b.ReportAllocs() + for b.Loop() { + bl := NewBuilder(nil, BuilderOptions{Fields: fields}) + for i := range 1000 { + r := bl.Row() + r.SetEpoch(epoch.Epoch(i / len(hosts))) + r.SetTag(0, hosts[i%len(hosts)]) + r.SetTag(1, dc) + r.AddValue(nil) + if err := r.Commit(); err != nil { + b.Fatal(err) + } + } + if _, err := bl.Finish(); err != nil { + b.Fatal(err) + } + } +} + +func TestBuilderEquivalentTagsShareSeries(t *testing.T) { + b := NewBuilder(nil, BuilderOptions{ + Fields: testBuilderFields(), + TagString: func(_ timeseries.FieldDefinition, raw []byte) string { return strings.ToLower(string(raw)) }, + }) + commitRows(t, b, + testRow{e: 1, host: "a", dc: "x", v: 1.0}, + testRow{e: 2, host: "A", dc: "X", v: 2.0}, + testRow{e: 3, host: "A", dc: "X", v: 3.0}, + testRow{e: 1, host: "b", dc: "x", v: 4.0}, + ) + // the second spelling is remembered, so its later rows skip the conversion + require.Len(t, b.results[0].lookup, 3) + ds, err := b.Finish() + require.NoError(t, err) + sl := ds.Results[0].SeriesList + require.Len(t, sl, 2) + require.Equal(t, Tags{"host": "a", "dc": "x"}, sl[0].Header.Tags) + require.Equal(t, []any{1.0, 2.0, 3.0}, pointValues(sl[0])) + requireSizes(t, sl[0]) + require.Equal(t, []any{4.0}, pointValues(sl[1])) +} + +func TestBuilderDuplicateTagNamesShareSeries(t *testing.T) { + fields := testBuilderFields() + fields.Tags[1].Name = fields.Tags[0].Name + b := NewBuilder(nil, BuilderOptions{Fields: fields}) + commitRows(t, b, + testRow{e: 1, host: "a", dc: "x", v: 1.0}, + testRow{e: 2, host: "b", dc: "x", v: 2.0}, + ) + ds, err := b.Finish() + require.NoError(t, err) + sl := ds.Results[0].SeriesList + require.Len(t, sl, 1) + require.Equal(t, Tags{"host": "x"}, sl[0].Header.Tags) + require.Equal(t, []epoch.Epoch{1, 2}, pointEpochs(sl[0])) +} + +func TestBuilderStartSeriesReopensIdenticalHeader(t *testing.T) { + value := timeseries.FieldDefinition{Name: "v", DataType: timeseries.String} + h := SeriesHeader{Name: "up", Tags: Tags{"job": "a"}, ValueFieldsList: timeseries.FieldDefinitions{value}} + b := NewBuilder(nil, BuilderOptions{}) + b.StartSeries(h) + require.NoError(t, b.AppendPoint(Point{Epoch: 2, Values: []any{"2"}})) + b.StartSeries(SeriesHeader{Name: "up", Tags: Tags{"job": "b"}, ValueFieldsList: h.ValueFieldsList}) + require.NoError(t, b.AppendPoint(Point{Epoch: 1, Values: []any{"1"}})) + // attributes outside the header's identity, such as output positions, do not split a series + moved := value + moved.OutputPosition = 3 + b.StartSeries(SeriesHeader{Name: "up", Tags: Tags{"job": "a"}, ValueFieldsList: timeseries.FieldDefinitions{moved}}) + require.NoError(t, b.AppendPoint(Point{Epoch: 1, Values: []any{"1"}})) + ds, err := b.Finish() + require.NoError(t, err) + sl := ds.Results[0].SeriesList + require.Len(t, sl, 2) + require.Equal(t, []epoch.Epoch{1, 2}, pointEpochs(sl[0])) + require.Zero(t, sl[0].Header.ValueFieldsList[0].OutputPosition) + require.Equal(t, []epoch.Epoch{1}, pointEpochs(sl[1])) +} + +func TestSameSeries(t *testing.T) { + fd := timeseries.FieldDefinition{Name: "v", DataType: timeseries.Float64} + base := func() SeriesHeader { + return SeriesHeader{ + Name: "n", QueryStatement: "q", Tags: Tags{"k": "v"}, TimestampField: fd, + ValueFieldsList: timeseries.FieldDefinitions{fd}, UntrackedFieldsList: timeseries.FieldDefinitions{fd}, + } + } + a := base() + b := base() + require.True(t, sameSeries(&a, &b)) + b.TagFieldsList = timeseries.FieldDefinitions{fd} + b.Size = 9 + require.True(t, sameSeries(&a, &b)) + for name, mutate := range map[string]func(*SeriesHeader){ + "name": func(h *SeriesHeader) { h.Name = "x" }, + "query": func(h *SeriesHeader) { h.QueryStatement = "x" }, + "tags": func(h *SeriesHeader) { h.Tags = Tags{"k": "x"} }, + "timestamp": func(h *SeriesHeader) { h.TimestampField.DataType = timeseries.Int64 }, + "values": func(h *SeriesHeader) { h.ValueFieldsList = nil }, + "value": func(h *SeriesHeader) { h.ValueFieldsList[0].Name = "x" }, + "untracked": func(h *SeriesHeader) { h.UntrackedFieldsList = nil }, + } { + b := base() + mutate(&b) + require.False(t, sameSeries(&a, &b), name) + } +} diff --git a/pkg/timeseries/dataset/dataset.go b/pkg/timeseries/dataset/dataset.go index 1e71e542f..6b43892bd 100644 --- a/pkg/timeseries/dataset/dataset.go +++ b/pkg/timeseries/dataset/dataset.go @@ -662,14 +662,14 @@ func (ds *DataSet) DefaultRangeCropper(e timeseries.Extent) { if start < l && end <= l && end > start { s.Points = s.Points.CloneRange(start, end) s.PointSize = s.Points.Size() + sl[index] = s } - sl[index] = s return nil }) j++ } eg.Wait() - ds.Results[i].SeriesList = sl[:j] + ds.Results[i].SeriesList = slices.DeleteFunc(sl[:j], func(s *Series) bool { return s == nil }) } } diff --git a/pkg/timeseries/dataset/dataset_test.go b/pkg/timeseries/dataset/dataset_test.go index 6f3aec300..a947aa40c 100644 --- a/pkg/timeseries/dataset/dataset_test.go +++ b/pkg/timeseries/dataset/dataset_test.go @@ -18,7 +18,6 @@ package dataset import ( "fmt" - "math/rand" "testing" "time" @@ -26,6 +25,7 @@ import ( "github.com/trickstercache/trickster/v2/pkg/timeseries" "github.com/trickstercache/trickster/v2/pkg/timeseries/epoch" "github.com/trickstercache/trickster/v2/pkg/timeseries/merge" + "github.com/trickstercache/trickster/v2/pkg/util/weak/weaktest" ) func testDataSet() *DataSet { @@ -42,6 +42,24 @@ func testDataSet() *DataSet { return ds } +func TestCropToRangeDropsSeriesWithoutMatchingPoints(t *testing.T) { + for _, bounds := range [][2]int64{{2, 3}, {5, 6}, {9, 10}} { + t.Run(fmt.Sprint(bounds), func(t *testing.T) { + ds := &DataSet{ + ExtentList: timeseries.ExtentList{{Start: time.Unix(0, 0), End: time.Unix(12, 0)}}, + Results: Results{&Result{SeriesList: SeriesList{ + &Series{Points: Points{{Epoch: epoch.Epoch(time.Unix(4, 0).UnixNano()), Values: []any{1}}, {Epoch: epoch.Epoch(time.Unix(8, 0).UnixNano()), Values: []any{2}}}}, + &Series{Points: Points{{Epoch: epoch.Epoch(time.Unix(bounds[0], 0).UnixNano()), Values: []any{3}}}}, + }}}, + } + ds.CropToRange(timeseries.Extent{Start: time.Unix(bounds[0], 0), End: time.Unix(bounds[1], 0)}) + if ds.SeriesCount() != 1 || ds.ValueCount() != 1 || ds.Results[0].SeriesList[0].Points[0].Values[0] != 3 { + t.Fatalf("crop retained points from a disjoint series: %+v", ds.Results[0].SeriesList) + } + }) + } +} + func genTestDataSet(seriesCount int, resultsCount int) *DataSet { if resultsCount <= 0 { logger.Error("resultsCount must be greater than 0", nil) @@ -726,7 +744,7 @@ func genBenchmarkPoint(e epoch.Epoch, valuect int) Point { Values: make([]any, valuect), } for i := range valuect { - out.Values[i] = rand.Int() % 1000 + out.Values[i] = weaktest.IntN(1000) } return out } diff --git a/pkg/timeseries/dataset/series_header.go b/pkg/timeseries/dataset/series_header.go index c0991ebbc..da10ddd53 100644 --- a/pkg/timeseries/dataset/series_header.go +++ b/pkg/timeseries/dataset/series_header.go @@ -20,7 +20,9 @@ package dataset import ( "fmt" + "maps" "slices" + "strconv" "strings" "github.com/trickstercache/trickster/v2/pkg/checksum/fnv" @@ -57,26 +59,57 @@ type SeriesHeader struct { // seriesHeaderFNVHash is the shared FNV payload for CalculateHash; queryStatement // is the string mixed into the hash in the same position as QueryStatement. func seriesHeaderFNVHash(sh *SeriesHeader, queryStatement string) Hash { + // strings are length-prefixed and lists counted, so no two headers share an input; + // sameSeries must compare exactly these attributes h := fnv.NewInlineFNV64a() - h.Write([]byte(sh.Name)) - h.Write([]byte(queryStatement)) + hashString(&h, sh.Name) + hashString(&h, queryStatement) + hashCount(&h, len(sh.Tags)) for _, k := range sh.Tags.Keys() { - h.Write([]byte(k)) - h.Write([]byte(sh.Tags[k])) + hashString(&h, k) + hashString(&h, sh.Tags[k]) } - for _, fd := range sh.ValueFieldsList { - h.Write([]byte(fd.Name)) - h.Write([]byte{byte(fd.DataType)}) - } - for _, fd := range sh.UntrackedFieldsList { - h.Write([]byte(fd.Name)) - h.Write([]byte{byte(fd.DataType)}) - } - h.Write([]byte(sh.TimestampField.Name)) - h.Write([]byte{byte(sh.TimestampField.DataType)}) + hashFields(&h, sh.ValueFieldsList) + hashFields(&h, sh.UntrackedFieldsList) + hashField(&h, sh.TimestampField) return Hash(h.Sum64()) } +func hashString(h *fnv.InlineFNV64a, s string) { + var buf [24]byte + _, _ = h.Write(append(strconv.AppendInt(buf[:0], int64(len(s)), 10), ':')) + _, _ = h.WriteString(s) +} + +func hashCount(h *fnv.InlineFNV64a, n int) { + var buf [24]byte + _, _ = h.Write(append(strconv.AppendInt(buf[:0], int64(n), 10), '#')) +} + +func hashField(h *fnv.InlineFNV64a, fd timeseries.FieldDefinition) { + hashString(h, fd.Name) + _, _ = h.Write([]byte{byte(fd.DataType)}) +} + +func hashFields(h *fnv.InlineFNV64a, fds timeseries.FieldDefinitions) { + hashCount(h, len(fds)) + for _, fd := range fds { + hashField(h, fd) + } +} + +func sameSeries(a, b *SeriesHeader) bool { + // compares exactly the attributes that seriesHeaderFNVHash covers + return a.Name == b.Name && a.QueryStatement == b.QueryStatement && + maps.Equal(a.Tags, b.Tags) && sameField(a.TimestampField, b.TimestampField) && + slices.EqualFunc(a.ValueFieldsList, b.ValueFieldsList, sameField) && + slices.EqualFunc(a.UntrackedFieldsList, b.UntrackedFieldsList, sameField) +} + +func sameField(a, b timeseries.FieldDefinition) bool { + return a.Name == b.Name && a.DataType == b.DataType +} + // CalculateHash sums the FNV64a hash for the Header and stores it to the Hash member func (sh *SeriesHeader) CalculateHash(rehash ...bool) Hash { if (len(rehash) == 0 || !rehash[0]) && sh.hash > 0 { diff --git a/pkg/timeseries/dataset/series_identity_test.go b/pkg/timeseries/dataset/series_identity_test.go new file mode 100644 index 000000000..aff9c9867 --- /dev/null +++ b/pkg/timeseries/dataset/series_identity_test.go @@ -0,0 +1,142 @@ +/* + * Copyright 2018 The Trickster Authors + * + * 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 dataset + +import ( + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/timeseries" + "github.com/trickstercache/trickster/v2/pkg/timeseries/epoch" + "github.com/trickstercache/trickster/v2/pkg/timeseries/merge" + + "github.com/stretchr/testify/require" +) + +func TestSeriesHeaderHashIsUnambiguous(t *testing.T) { + field := func(name string, dt timeseries.FieldDataType) timeseries.FieldDefinitions { + return timeseries.FieldDefinitions{{Name: name, DataType: dt}} + } + // each pair once produced identical hash input by running adjacent strings together + pairs := map[string][2]SeriesHeader{ + "tag boundary": {{Tags: Tags{"ab": "c"}}, {Tags: Tags{"a": "bc"}}}, + "name and query": {{Name: "ab", QueryStatement: "c"}, {Name: "a", QueryStatement: "bc"}}, + "tags and values": { + {Tags: Tags{"x": "y\x01"}}, + {Tags: Tags{"x": "y"}, ValueFieldsList: field("", 1)}, + }, + "values and untracked": { + {ValueFieldsList: field("a", 1)}, + {UntrackedFieldsList: field("a", 1)}, + }, + } + for name, pair := range pairs { + a, b := pair[0], pair[1] + require.NotEqual(t, a.CalculateHash(), b.CalculateHash(), name) + require.False(t, sameSeries(&a, &b), name) + } + h := SeriesHeader{Name: "n", Tags: Tags{"k": "v"}} + require.Equal(t, h.CalculateHash(), h.CalculateHashWithQueryStatement("")) +} + +func collidingSeries(host string, epochs ...epoch.Epoch) *Series { + s := &Series{Header: SeriesHeader{Name: "s", Tags: Tags{"host": host}}} + s.Header.hash = 7 // every series here shares one hash + for _, e := range epochs { + s.Points = append(s.Points, Point{Epoch: e, Size: 16, Values: []any{int64(e)}}) + } + s.PointSize = s.Points.Size() + return s +} + +func epochsByHost(sl SeriesList) map[string][]epoch.Epoch { + out := make(map[string][]epoch.Epoch, len(sl)) + for _, s := range sl { + out[s.Header.Tags["host"]] = pointEpochs(s) + } + return out +} + +func TestSeriesIndex(t *testing.T) { + idx := newSeriesIndex(0, headerOf) + a, b, c := collidingSeries("a"), collidingSeries("b"), collidingSeries("c") + idx.add(7, a) + idx.add(7, b) + for _, s := range []*Series{a, b} { + got, ok := idx.find(7, &collidingSeries(s.Header.Tags["host"]).Header) + require.True(t, ok) + require.Same(t, s, got) + } + _, ok := idx.find(7, &c.Header) + require.False(t, ok) + _, ok = idx.find(8, &a.Header) + require.False(t, ok) + idx.reset() + _, ok = idx.find(7, &a.Header) + require.False(t, ok) + require.Empty(t, idx.spill) +} + +func TestMergesKeepCollidingSeriesApart(t *testing.T) { + lists := func() (SeriesList, SeriesList) { + // the receiver repeats "b" and the incoming list repeats "c"; repeats are dropped + return SeriesList{collidingSeries("a", 1), collidingSeries("b", 1), collidingSeries("b", 9)}, + SeriesList{collidingSeries("a", 2), collidingSeries("c", 3), collidingSeries("c", 4)} + } + want := map[string][]epoch.Epoch{"a": {1, 2}, "b": {1}, "c": {3}} + + sl, sl2 := lists() + require.Equal(t, want, epochsByHost(sl.Merge(sl2, true))) + + sl, sl2 = lists() + opts := MergeOpts{SortPoints: true, Strategy: merge.StrategySum} + require.Equal(t, want, epochsByHost(sl.MergeWithOpts(sl2, opts))) + + sl, sl2 = lists() + out := sl.mergeCollection([]SeriesList{sl2, {collidingSeries("b", 5), collidingSeries("d", 6)}}, + MergeOpts{SortPoints: true}) + require.Equal(t, map[string][]epoch.Epoch{"a": {1, 2}, "b": {1, 5}, "c": {3}, "d": {6}}, + epochsByHost(out)) + + // nil series are skipped, and an empty receiver starts from the first member + var empty SeriesList + out = empty.mergeCollection([]SeriesList{{collidingSeries("a", 1), nil}, {nil, collidingSeries("b", 2)}}, + MergeOpts{SortPoints: true}) + require.Equal(t, map[string][]epoch.Epoch{"a": {1}, "b": {2}}, epochsByHost(out)) + out = SeriesList{nil, collidingSeries("a", 1)}.mergeCollection( + []SeriesList{{collidingSeries("a", 2)}, {collidingSeries("b", 3)}}, MergeOpts{SortPoints: true}) + require.Equal(t, map[string][]epoch.Epoch{"a": {1, 2}, "b": {3}}, epochsByHost(out)) + require.Empty(t, empty.mergeCollection([]SeriesList{{}, nil}, MergeOpts{})) +} + +func TestDataSetMergeKeepsCollidingSeries(t *testing.T) { + trq := ×eries.TimeRangeQuery{Step: time.Second} + dataSet := func(series ...*Series) *DataSet { + return &DataSet{TimeRangeQuery: trq, Results: Results{{SeriesList: series}}} + } + ds := dataSet(collidingSeries("a", 1), collidingSeries("b", 1)) + ds.Merge(true, dataSet(collidingSeries("a", 2), collidingSeries("c", 3)), + dataSet(collidingSeries("b", 4))) + require.Equal(t, map[string][]epoch.Epoch{"a": {1, 2}, "b": {1, 4}, "c": {3}}, + epochsByHost(ds.Results[0].SeriesList)) +} + +func TestEqualHeaderDetectsCollision(t *testing.T) { + a := SeriesList{collidingSeries("a")} + require.True(t, a.EqualHeader(SeriesList{collidingSeries("a")})) + require.False(t, a.EqualHeader(SeriesList{collidingSeries("b")})) +} diff --git a/pkg/timeseries/dataset/series_index.go b/pkg/timeseries/dataset/series_index.go new file mode 100644 index 000000000..e10ec1e47 --- /dev/null +++ b/pkg/timeseries/dataset/series_index.go @@ -0,0 +1,64 @@ +/* + * Copyright 2018 The Trickster Authors + * + * 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 dataset + +type seriesIndex[T any] struct { + header func(T) *SeriesHeader + first map[Hash]T + spill map[Hash][]T +} + +func newSeriesIndex[T any](size int, header func(T) *SeriesHeader) *seriesIndex[T] { + return &seriesIndex[T]{header: header, first: make(map[Hash]T, size)} +} + +func (x *seriesIndex[T]) find(h Hash, sh *SeriesHeader) (T, bool) { + // the hash narrows the search and the header comparison decides it, so series + // whose hashes collide are never treated as one + if v, ok := x.first[h]; ok { + if sameSeries(x.header(v), sh) { + return v, true + } + for _, v := range x.spill[h] { + if sameSeries(x.header(v), sh) { + return v, true + } + } + } + var zero T + return zero, false +} + +func (x *seriesIndex[T]) add(h Hash, v T) { + if _, ok := x.first[h]; !ok { + x.first[h] = v + return + } + if x.spill == nil { + x.spill = make(map[Hash][]T) + } + x.spill[h] = append(x.spill[h], v) +} + +func (x *seriesIndex[T]) reset() { + clear(x.first) + clear(x.spill) +} + +func headerOf(s *Series) *SeriesHeader { + return &s.Header +} diff --git a/pkg/timeseries/dataset/series_list.go b/pkg/timeseries/dataset/series_list.go index 01a916f40..c44e78bb2 100644 --- a/pkg/timeseries/dataset/series_list.go +++ b/pkg/timeseries/dataset/series_list.go @@ -24,7 +24,6 @@ import ( "sync" "github.com/trickstercache/trickster/v2/pkg/timeseries/merge" - "github.com/trickstercache/trickster/v2/pkg/util/sets" "golang.org/x/sync/errgroup" ) @@ -37,16 +36,20 @@ type SeriesList []*Series // Merge merges sl2 into the subject SeriesList, using sl2's authoritative order // to adaptively reorder the existing+merged list such that it best emulates // the fully constituted series order as it would be served by the origin. -// Merge assumes that a *Series in both lists, having the identical header hash, -// are the same series and will merge sl2[i].Points into sl.Points +// Merge treats a *Series in both lists with identical headers as the same +// series and merges its Points from sl2 into those from sl. func (sl SeriesList) Merge(sl2 SeriesList, sortPoints bool) SeriesList { + return sl.merge(sl2, func(p, p2 Points) Points { return MergePoints(p, p2, sortPoints) }) +} + +func (sl SeriesList) merge(sl2 SeriesList, mergePoints func(p, p2 Points) Points) SeriesList { if len(sl2) == 0 { return sl.Clone() } if len(sl) == 0 { return sl2.Clone() } - m := make(map[Hash]*Series, len(sl)+len(sl2)) + idx := newSeriesIndex(len(sl)+len(sl2), headerOf) out := make(SeriesList, len(sl)+len(sl2)) var k int for _, s := range sl { @@ -54,33 +57,33 @@ func (sl SeriesList) Merge(sl2 SeriesList, sortPoints bool) SeriesList { continue } h := s.Header.CalculateHash() - if _, ok := m[h]; ok { + if _, ok := idx.find(h, &s.Header); ok { continue } out[k] = s - m[h] = s + idx.add(h, s) k++ } - seen := make(sets.Set[Hash], len(sl2)) + seen := newSeriesIndex(len(sl2), headerOf) var wg sync.WaitGroup for _, s := range sl2 { if s == nil { continue } h := s.Header.CalculateHash() - if seen.Contains(h) { + if _, ok := seen.find(h, &s.Header); ok { continue } - seen.Set(h) - if cs, ok := m[h]; !ok { + seen.add(h, s) + if cs, ok := idx.find(h, &s.Header); !ok { // this series does not exist in sl1; add it into out out[k] = s - m[h] = s + idx.add(h, s) k++ } else { // series is in both sl1 and sl2; merge their points wg.Go(func() { - cs.Points = MergePoints(cs.Points, s.Points, sortPoints) + cs.Points = mergePoints(cs.Points, s.Points) cs.PointSize = cs.Points.Size() }) } @@ -104,7 +107,8 @@ func (sl SeriesList) EqualHeader(sl2 SeriesList) bool { if v == nil || sl2[i] == nil { return false } - if v.Header.CalculateHash() != sl2[i].Header.CalculateHash() { + if v.Header.CalculateHash() != sl2[i].Header.CalculateHash() || + !sameSeries(&v.Header, &sl2[i].Header) { return false } } @@ -148,53 +152,7 @@ func (sl SeriesList) MergeWithOpts(sl2 SeriesList, opts MergeOpts) SeriesList { // fast path: legacy exact-match dedup return sl.Merge(sl2, opts.SortPoints) } - if len(sl2) == 0 { - return sl.Clone() - } - if len(sl) == 0 { - return sl2.Clone() - } - m := make(map[Hash]*Series, len(sl)+len(sl2)) - out := make(SeriesList, len(sl)+len(sl2)) - var k int - for _, s := range sl { - if s == nil { - continue - } - h := s.Header.CalculateHash() - if _, ok := m[h]; ok { - continue - } - out[k] = s - m[h] = s - k++ - } - seen := make(sets.Set[Hash], len(sl2)) - var wg sync.WaitGroup - for _, s := range sl2 { - if s == nil { - continue - } - h := s.Header.CalculateHash() - if seen.Contains(h) { - continue - } - seen.Set(h) - if cs, ok := m[h]; !ok { - out[k] = s - m[h] = s - k++ - } else { - wg.Go(func() { - cs.Points = MergePointsWithOpts(cs.Points, s.Points, opts) - cs.PointSize = cs.Points.Size() - }) - } - } - wg.Wait() - out = out[:k] - out.SortByTags() - return out + return sl.merge(sl2, func(p, p2 Points) Points { return MergePointsWithOpts(p, p2, opts) }) } // mergeCollection merges several member lists while preserving the same @@ -229,18 +187,18 @@ func (sl SeriesList) mergeCollection(collection []SeriesList, opts MergeOpts) Se total += len(next) } out := make(SeriesList, total) - seriesByHash := make(map[Hash]*Series, total) + idx := newSeriesIndex(total, headerOf) var k int for _, s := range sl { if s == nil { continue } h := s.Header.CalculateHash() - if _, ok := seriesByHash[h]; ok { + if _, ok := idx.find(h, &s.Header); ok { continue } out[k] = s - seriesByHash[h] = s + idx.add(h, s) k++ } @@ -249,30 +207,30 @@ func (sl SeriesList) mergeCollection(collection []SeriesList, opts MergeOpts) Se series []*Series } jobs := make([]mergeJob, 0) - jobByHash := make(map[Hash]int) - seen := make(sets.Set[Hash]) + jobByTarget := make(map[*Series]int) + seen := newSeriesIndex(0, headerOf) for _, next := range nonEmpty { - clear(seen) + seen.reset() for _, s := range next { if s == nil { continue } h := s.Header.CalculateHash() - if seen.Contains(h) { + if _, ok := seen.find(h, &s.Header); ok { continue } - seen.Set(h) - target, ok := seriesByHash[h] + seen.add(h, s) + target, ok := idx.find(h, &s.Header) if !ok { out[k] = s - seriesByHash[h] = s + idx.add(h, s) k++ continue } - jobIndex, ok := jobByHash[h] + jobIndex, ok := jobByTarget[target] if !ok { jobIndex = len(jobs) - jobByHash[h] = jobIndex + jobByTarget[target] = jobIndex jobs = append(jobs, mergeJob{target: target}) } jobs[jobIndex].series = append(jobs[jobIndex].series, s) diff --git a/pkg/timeseries/dataset/series_test.go b/pkg/timeseries/dataset/series_test.go index 3c65c8e22..d2c86ba9b 100644 --- a/pkg/timeseries/dataset/series_test.go +++ b/pkg/timeseries/dataset/series_test.go @@ -62,7 +62,7 @@ func TestString(t *testing.T) { if s.String() != expected { t.Errorf("expected %s got %s", expected, s.String()) } - expected = "[16450490800955907542]" + expected = "[1032707601692489584]" sl := SeriesList{s} if sl.String() != expected { t.Errorf("expected %s got %s", expected, sl.String()) diff --git a/pkg/timeseries/dataset/stream/conformance_test.go b/pkg/timeseries/dataset/stream/conformance_test.go new file mode 100644 index 000000000..4f285e6ed --- /dev/null +++ b/pkg/timeseries/dataset/stream/conformance_test.go @@ -0,0 +1,260 @@ +/* + * Copyright 2018 The Trickster Authors + * + * 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 stream_test + +import ( + "cmp" + "encoding/json" + "io" + "maps" + "math" + "slices" + "strings" + "testing" + + "github.com/trickstercache/trickster/v2/pkg/timeseries" + "github.com/trickstercache/trickster/v2/pkg/timeseries/dataset" + "github.com/trickstercache/trickster/v2/pkg/timeseries/dataset/stream" + "github.com/trickstercache/trickster/v2/pkg/timeseries/dataset/stream/streamtest" + "github.com/trickstercache/trickster/v2/pkg/timeseries/epoch" + "github.com/trickstercache/trickster/v2/pkg/util/weak/weaktest" + + "github.com/stretchr/testify/require" +) + +func wantSeries(name string, tags dataset.Tags, fields timeseries.SeriesFields, + points ...dataset.Point, +) *dataset.Series { + return &dataset.Series{ + Header: dataset.SeriesHeader{ + Name: name, + Tags: tags, + TimestampField: fields.Timestamp, + TagFieldsList: fields.Tags, + ValueFieldsList: fields.Values, + }, + Points: points, + } +} + +func wantDataSet(series ...*dataset.Series) *dataset.DataSet { + return &dataset.DataSet{ + ExtentList: timeseries.ExtentList{testTRQ.Extent}, + Results: dataset.Results{{SeriesList: series}}, + } +} + +func pt(ms int64, v any) dataset.Point { + return dataset.Point{Epoch: epoch.Epoch(ms * 1e6), Values: []any{v}} +} + +func TestTSVConformance(t *testing.T) { + body := "time\thost\tvalue\r\n2000\ta\t3\n1000\ta\t1.5\n1000\tb\t2\n\n" + want := wantDataSet( + wantSeries("tsv", dataset.Tags{"host": "a"}, rowFields, pt(1000, 1.5), pt(2000, 3.0)), + wantSeries("tsv", dataset.Tags{"host": "b"}, rowFields, pt(1000, 2.0)), + ) + for _, s := range want.Results[0].SeriesList { + s.Header.QueryStatement = "q" + } + // the blank line is a row with the wrong column count + streamtest.Conformance(t, newTSVDecoder, streamtest.Case{ + TRQ: testTRQ, Body: []byte(body), WantErr: timeseries.ErrInvalidBody, + }) + streamtest.Conformance(t, newTSVDecoder, streamtest.Case{ + TRQ: testTRQ, Body: []byte(strings.TrimSuffix(body, "\n\n")), Want: want, + Shuffle: streamtest.ShuffleLines(1), + }) +} + +func TestTSVConformanceErrors(t *testing.T) { + tests := []struct { + name string + body string + err error + }{ + {"empty", "", errBadHeader}, + {"header", "time\tvalue\n1\t1\n", errBadHeader}, + {"columns", "time\thost\tvalue\n1\ta\n", timeseries.ErrInvalidBody}, + {"epoch", "time\thost\tvalue\nx\ta\t1\n", timeseries.ErrInvalidTimeFormat}, + {"value", "time\thost\tvalue\n1\ta\tx\n", stream.ErrInvalidValue}, + {"duplicate", "time\thost\tvalue\n1\ta\t1\n1\ta\t2\n", dataset.ErrDuplicateEpoch}, + {"unordered duplicate", "time\thost\tvalue\n2\ta\t1\n1\ta\t1\n2\ta\t2\n", dataset.ErrDuplicateEpoch}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + streamtest.Conformance(t, newTSVDecoder, streamtest.Case{ + TRQ: testTRQ, Body: []byte(test.body), WantErr: test.err, + }) + }) + } +} + +func TestMatrixConformance(t *testing.T) { + body := `{"status":"success","data":{"resultType":"matrix","result":[ + {"metric":{"job":"a"},"values":[[1,"1"],[2.5,"2.5"]]}, + {"metric":{"job":"b"},"values":[[2,"NaN"],[1,"+Inf"]],"extra":[{}]}]},"warnings":[]} + ` + fields := timeseries.SeriesFields{Values: timeseries.FieldDefinitions{ + {Name: "value", DataType: timeseries.Float64, Role: timeseries.RoleValue}}} + want := wantDataSet( + wantSeries("matrix", dataset.Tags{"job": "a"}, fields, pt(1000, 1.0), pt(2500, 2.5)), + wantSeries("matrix", dataset.Tags{"job": "b"}, fields, pt(1000, math.Inf(1)), pt(2000, math.NaN())), + ) + want.Status = "success" + streamtest.Conformance(t, newMatrixDecoder, streamtest.Case{TRQ: testTRQ, Body: []byte(body), Want: want}) +} + +func TestMatrixConformanceErrors(t *testing.T) { + tests := []struct { + name string + body string + err error + }{ + {"empty", "", streamtest.ErrAny}, + {"truncated", `{"data":{"result":[{"metric":{}`, streamtest.ErrAny}, + {"null result", `{"data":{"result":null}}`, stream.ErrNull}, + {"not an object", `[]`, stream.ErrUnexpectedToken}, + {"trailing", `{} {}`, stream.ErrTrailingData}, + {"trailing garbage", `{} x`, streamtest.ErrAny}, + {"values first", `{"data":{"result":[{"values":[],"metric":{}}]}}`, errMetricFirst}, + {"bad pair", `{"data":{"result":[{"metric":{},"values":[[1]]}]}}`, timeseries.ErrInvalidBody}, + {"bad time", `{"data":{"result":[{"metric":{},"values":[["x","1"]]}]}}`, timeseries.ErrInvalidTimeFormat}, + {"bad value", `{"data":{"result":[{"metric":{},"values":[[1,"x"]]}]}}`, stream.ErrInvalidValue}, + {"duplicate", `{"data":{"result":[{"metric":{},"values":[[1,"1"],[1,"2"]]}]}}`, dataset.ErrDuplicateEpoch}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + streamtest.Conformance(t, newMatrixDecoder, streamtest.Case{ + TRQ: testTRQ, Body: []byte(test.body), WantErr: test.err, + }) + }) + } +} + +func shuffleRows(body []byte, rng *weaktest.Rand) []byte { + var doc map[string]json.RawMessage + var rows []json.RawMessage + if json.Unmarshal(body, &doc) != nil || json.Unmarshal(doc["rows"], &rows) != nil { + return body + } + rng.Shuffle(len(rows), func(i, j int) { rows[i], rows[j] = rows[j], rows[i] }) + doc["rows"], _ = json.Marshal(rows) + out, _ := json.Marshal(doc) + return out +} + +func legacyRows(r io.Reader, trq *timeseries.TimeRangeQuery) (timeseries.Timeseries, error) { + // decodes the whole document before building its DataSet, like pre-stream unmarshalers + var doc struct { + Rows [][]any `json:"rows"` + } + dec := json.NewDecoder(r) + dec.UseNumber() + if err := dec.Decode(&doc); err != nil { + return nil, err + } + byHost := map[string]*dataset.Series{} + for _, row := range doc.Rows { + host := row[1].(string) + s, ok := byHost[host] + if !ok { + s = wantSeries("rows", dataset.Tags{"host": host}, rowFields) + byHost[host] = s + } + ms, _ := row[0].(json.Number).Int64() + var v any + if n, ok := row[2].(json.Number); ok { + v, _ = n.Float64() + } else if str, ok := row[2].(string); ok { + v, _ = json.Number(str).Float64() + } + s.Points = append(s.Points, pt(ms, v)) + } + sl := make(dataset.SeriesList, 0, len(byHost)) + for _, host := range slices.Sorted(maps.Keys(byHost)) { + s := byHost[host] + slices.SortStableFunc(s.Points, func(a, b dataset.Point) int { return cmp.Compare(a.Epoch, b.Epoch) }) + sl = append(sl, s) + } + ds := wantDataSet(sl...) + ds.TimeRangeQuery = trq + return ds, nil +} + +func TestRowsConformance(t *testing.T) { + // "\u0061" spells the same tag value as "a", so both rows belong to one series + body := `{"rows":[[2000,"b",1],[1000,"a",null],[1000,"b","2.5"],[3000,"c\"d",-4e-1],[4000,"\u0061",5]], + "total":5,"extra":{"x":[1,2,{"y":null}]}}` + want := wantDataSet( + wantSeries("rows", dataset.Tags{"host": "a"}, rowFields, pt(1000, nil), pt(4000, 5.0)), + wantSeries("rows", dataset.Tags{"host": "b"}, rowFields, pt(1000, 2.5), pt(2000, 1.0)), + wantSeries("rows", dataset.Tags{"host": `c"d`}, rowFields, pt(3000, -0.4)), + ) + streamtest.Conformance(t, newRowsDecoder, streamtest.Case{ + TRQ: testTRQ, Body: []byte(body), Want: want, Shuffle: shuffleRows, Legacy: legacyRows, + }) + streamtest.Conformance(t, newRowsDecoder, streamtest.Case{ + TRQ: testTRQ, Body: []byte(`{"rows":[[1,"a",1]],"total":2}`), WantErr: timeseries.ErrInvalidBody, + }) +} + +func TestUnmarshalerAdapters(t *testing.T) { + u := stream.ReaderUnmarshaler(newTSVDecoder) + _, err := u(nil, testTRQ) + require.ErrorIs(t, err, timeseries.ErrInvalidBody) + + ts, err := stream.BytesUnmarshaler(newTSVDecoder)([]byte("time\thost\tvalue\n1\ta\t1\n"), testTRQ) + require.NoError(t, err) + require.Equal(t, int64(1), ts.ValueCount()) + + errNew := io.ErrClosedPipe + failing := func(*timeseries.TimeRangeQuery) (stream.Decoder, error) { return nil, errNew } + _, err = stream.BytesUnmarshaler(failing)(nil, testTRQ) + require.ErrorIs(t, err, errNew) + + nilResult := func(*timeseries.TimeRangeQuery) (stream.Decoder, error) { + return stream.NewLines(func([]byte) error { return nil }, + func() (timeseries.Timeseries, error) { return nil, nil }), nil + } + _, err = stream.BytesUnmarshaler(nilResult)([]byte("x"), testTRQ) + require.ErrorIs(t, err, timeseries.ErrInvalidBody) +} + +func BenchmarkRowsDecoders(b *testing.B) { + var sb strings.Builder + sb.WriteString(`{"rows":[`) + const rows = 5000 + for i := range rows { + if i > 0 { + sb.WriteByte(',') + } + sb.WriteString(`[`) + sb.WriteString(strings.Repeat("1", 1+i%3)) + sb.WriteString(`000,"host-`) + sb.WriteByte(byte('a' + i%8)) + sb.WriteString(`",1.5]`) + } + sb.WriteString(`],"total":5000}`) + body := []byte(sb.String()) + b.Run("legacy", func(b *testing.B) { + streamtest.Bench(b, legacyRows, testTRQ, body) + }) + b.Run("stream", func(b *testing.B) { + streamtest.Bench(b, stream.ReaderUnmarshaler(newRowsDecoder), testTRQ, body) + }) +} diff --git a/pkg/timeseries/dataset/stream/decoders_test.go b/pkg/timeseries/dataset/stream/decoders_test.go new file mode 100644 index 000000000..070c3b042 --- /dev/null +++ b/pkg/timeseries/dataset/stream/decoders_test.go @@ -0,0 +1,211 @@ +/* + * Copyright 2018 The Trickster Authors + * + * 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 stream_test + +import ( + "encoding/json" + "errors" + "time" + + "github.com/trickstercache/trickster/v2/pkg/timeseries" + "github.com/trickstercache/trickster/v2/pkg/timeseries/dataset" + "github.com/trickstercache/trickster/v2/pkg/timeseries/dataset/stream" + "github.com/trickstercache/trickster/v2/pkg/timeseries/epoch" +) + +// The decoders below show how providers use the stream package and Builder. + +var testTRQ = ×eries.TimeRangeQuery{ + Statement: "q", + Extent: timeseries.Extent{Start: time.Unix(0, 0), End: time.Unix(10, 0)}, + Step: time.Second, +} + +var rowFields = timeseries.SeriesFields{ + Timestamp: timeseries.FieldDefinition{Name: "time", DataType: timeseries.DateTimeUnixMilli, + Role: timeseries.RoleTimestamp}, + Tags: timeseries.FieldDefinitions{{Name: "host", DataType: timeseries.String, + Role: timeseries.RoleTag, OutputPosition: 1}}, + Values: timeseries.FieldDefinitions{{Name: "value", DataType: timeseries.Float64, + Role: timeseries.RoleValue, OutputPosition: 2}}, +} + +var ( + errBadHeader = errors.New("bad header") + errMetricFirst = errors.New("metric must precede values") +) + +func newTSVDecoder(trq *timeseries.TimeRangeQuery) (stream.Decoder, error) { + // decodes "time\thost\tvalue" rows that may arrive in any order + b := dataset.NewBuilder(trq, dataset.BuilderOptions{Fields: rowFields, SeriesName: "tsv", + QueryStatement: trq.Statement, Duplicates: dataset.DuplicatesError}) + var header bool + var cols [][]byte + onLine := func(line []byte) error { + if !header { + if string(line) != "time\thost\tvalue" { + return errBadHeader + } + header = true + return nil + } + cols = stream.SplitFields(line, '\t', cols) + if len(cols) != 3 { + return timeseries.ErrInvalidBody + } + ep, err := epoch.ParseDecimal(cols[0], timeseries.DateTimeUnixMilli) + if err != nil { + return err + } + v, err := stream.ParseValue(cols[2], timeseries.Float64) + if err != nil { + return err + } + r := b.Row() + r.SetEpoch(ep) + r.SetTag(0, cols[1]) + r.AddValue(v) + return r.Commit() + } + return stream.NewLines(onLine, func() (timeseries.Timeseries, error) { + if !header { + return nil, errBadHeader + } + return b.Finish() + }), nil +} + +func newMatrixDecoder(trq *timeseries.TimeRangeQuery) (stream.Decoder, error) { + // decodes a Prometheus-style matrix, one series at a time + b := dataset.NewBuilder(trq, dataset.BuilderOptions{Duplicates: dataset.DuplicatesError}) + valueFields := timeseries.FieldDefinitions{{Name: "value", DataType: timeseries.Float64, + Role: timeseries.RoleValue}} + var status string + var pair []json.RawMessage + series := func(dec *json.Decoder) error { + defer b.EndSeries() + var open bool + return stream.Object(dec, func(key string) error { + switch key { + case "metric": + var tags dataset.Tags + if err := dec.Decode(&tags); err != nil { + return err + } + b.StartSeries(dataset.SeriesHeader{Name: "matrix", Tags: tags, ValueFieldsList: valueFields}) + open = true + return nil + case "values": + if !open { + return errMetricFirst + } + return stream.Array(dec, func() error { + if err := dec.Decode(&pair); err != nil { + return err + } + if len(pair) != 2 { + return timeseries.ErrInvalidBody + } + ep, err := epoch.ParseDecimal(pair[0], timeseries.DateTimeUnixSecs) + if err != nil { + return err + } + v, err := stream.ParseJSONValue(pair[1], timeseries.Float64) + if err != nil { + return err + } + r := b.Row() + r.SetEpoch(ep) + r.AddValue(v) + return r.Commit() + }) + } + return stream.Skip(dec) + }) + } + walk := func(dec *json.Decoder) error { + return stream.Object(dec, func(key string) error { + switch key { + case "status": + return dec.Decode(&status) + case "data": + return stream.Object(dec, func(key string) error { + if key != "result" { + return stream.Skip(dec) + } + return stream.Array(dec, func() error { return series(dec) }) + }) + } + return stream.Skip(dec) + }) + } + return stream.NewJSON(walk, func() (timeseries.Timeseries, error) { + ds, err := b.Finish() + if err != nil { + return nil, err + } + ds.Status = status + return ds, nil + }), nil +} + +func newRowsDecoder(trq *timeseries.TimeRangeQuery) (stream.Decoder, error) { + // decodes {"rows":[[time,host,value],...],"total":n} rows that may arrive in + // any order, checking the trailing total once all rows are read + b := dataset.NewBuilder(trq, dataset.BuilderOptions{Fields: rowFields, SeriesName: "rows", + TagString: stream.JSONTagString, SortSeries: true}) + var row []json.RawMessage + var count, total int + walk := func(dec *json.Decoder) error { + return stream.Object(dec, func(key string) error { + switch key { + case "rows": + return stream.Array(dec, func() error { + if err := dec.Decode(&row); err != nil { + return err + } + if len(row) != 3 { + return timeseries.ErrInvalidBody + } + ep, err := epoch.ParseDecimal(row[0], timeseries.DateTimeUnixMilli) + if err != nil { + return err + } + v, err := stream.ParseJSONValue(row[2], timeseries.Float64) + if err != nil { + return err + } + r := b.Row() + r.SetEpoch(ep) + r.SetTag(0, row[1]) + r.AddValue(v) + count++ + return r.Commit() + }) + case "total": + return dec.Decode(&total) + } + return stream.Skip(dec) + }) + } + return stream.NewJSON(walk, func() (timeseries.Timeseries, error) { + if count != total { + return nil, timeseries.ErrInvalidBody + } + return b.Finish() + }), nil +} diff --git a/pkg/timeseries/dataset/stream/json.go b/pkg/timeseries/dataset/stream/json.go new file mode 100644 index 000000000..3ab92d81e --- /dev/null +++ b/pkg/timeseries/dataset/stream/json.go @@ -0,0 +1,256 @@ +/* + * Copyright 2018 The Trickster Authors + * + * 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 stream + +import ( + "bytes" + "encoding/json" + "errors" + "fmt" + "io" + + "github.com/trickstercache/trickster/v2/pkg/timeseries" +) + +var ( + // ErrNull is returned by Object and Array for a JSON null, which is consumed, + // so callers that accept a null may continue. + ErrNull = fmt.Errorf("%w: unexpected null JSON value", timeseries.ErrInvalidBody) + // ErrUnexpectedToken indicates a JSON value was not the expected object or array. + ErrUnexpectedToken = fmt.Errorf("%w: unexpected JSON token", timeseries.ErrInvalidBody) + // ErrTrailingData indicates input remained after the walk returned. + ErrTrailingData = fmt.Errorf("%w: unconsumed JSON input after document", timeseries.ErrInvalidBody) + // ErrValueNotConsumed indicates an Object or Array visitor returned without + // consuming the value it was called for. + ErrValueNotConsumed = errors.New("visitor did not consume a JSON value") +) + +// JSON is a Decoder for one JSON document that is walked token by token, so +// only the current token or decoded element is held in memory. +type JSON struct { + walk func(dec *json.Decoder) error + finish FinishFunc + buf []byte + read bool + err error + done bool +} + +var _ Decoder = (*JSON)(nil) + +// NewJSON returns a JSON decoder. walk must consume exactly one JSON value +// from dec, which has UseNumber set, and finish is called by Finish afterward. +func NewJSON(walk func(dec *json.Decoder) error, finish FinishFunc) *JSON { + return &JSON{walk: walk, finish: finish} +} + +// Write buffers p until ReadFrom or Finish walks the document. Use ReadFrom to +// decode while the input is being read. +func (j *JSON) Write(p []byte) (int, error) { + if err := j.check(); err != nil { + return 0, err + } + j.buf = append(j.buf, p...) + return len(p), nil +} + +// ReadFrom walks the document from any previously written bytes followed by r, +// reading r to EOF. No input may be provided afterward. +func (j *JSON) ReadFrom(r io.Reader) (int64, error) { + if err := j.check(); err != nil { + return 0, err + } + j.read = true + cr := &countingReader{r: r} + var src io.Reader = cr + if len(j.buf) > 0 { + src = io.MultiReader(bytes.NewReader(j.buf), cr) + } + err := j.run(src) + j.buf = nil + if err != nil { + j.err = err + } + return cr.n, err +} + +// Finish walks any buffered input not yet walked, then calls the finish function. +func (j *JSON) Finish() (timeseries.Timeseries, error) { + if j.done { + return nil, ErrFinished + } + j.done = true + if j.err != nil { + return nil, j.err + } + if !j.read { + err := j.run(bytes.NewReader(j.buf)) + j.buf = nil + if err != nil { + j.err = err + return nil, err + } + } + return j.finish() +} + +func (j *JSON) check() error { + switch { + case j.done: + return ErrFinished + case j.err != nil: + return j.err + case j.read: + return ErrInputConsumed + } + return nil +} + +func (j *JSON) run(r io.Reader) error { + dec := json.NewDecoder(r) + dec.UseNumber() + err := j.walk(dec) + if err == nil { + // only whitespace may follow the document + _, err = dec.Token() + if errors.Is(err, io.EOF) { + return nil + } + if err == nil { + return ErrTrailingData + } + } + if errors.Is(err, io.EOF) { + return io.ErrUnexpectedEOF + } + return err +} + +type countingReader struct { + r io.Reader + n int64 +} + +func (c *countingReader) Read(p []byte) (int, error) { + n, err := c.r.Read(p) + c.n += int64(n) + return n, err +} + +// Object consumes a JSON object from dec, calling fn with each key in arrival +// order. fn must consume the key's value, e.g. with dec.Decode, Object, Array or Skip. +func Object(dec *json.Decoder, fn func(key string) error) error { + if err := open(dec, '{'); err != nil { + return err + } + for dec.More() { + tok, err := dec.Token() + if err != nil { + return err + } + key, ok := tok.(string) + if !ok { + return ErrUnexpectedToken + } + offset := dec.InputOffset() + if err := fn(key); err != nil { + return err + } + if dec.InputOffset() == offset { + return ErrValueNotConsumed + } + } + return closeDelim(dec, '}') +} + +// Array consumes a JSON array from dec, calling fn once per element. fn must +// consume the element, e.g. with dec.Decode, Object, Array or Skip. +func Array(dec *json.Decoder, fn func() error) error { + if err := open(dec, '['); err != nil { + return err + } + for dec.More() { + offset := dec.InputOffset() + if err := fn(); err != nil { + return err + } + if dec.InputOffset() == offset { + return ErrValueNotConsumed + } + } + return closeDelim(dec, ']') +} + +// Skip consumes and discards the next JSON value from dec one token at a time, so a +// large skipped value is never held in memory. It fails if no value comes next. +func Skip(dec *json.Decoder) error { + var depth int + for { + tok, err := dec.Token() + if err != nil { + return err + } + switch tok { + case json.Delim('['), json.Delim('{'): + depth++ + case json.Delim(']'), json.Delim('}'): + depth-- + } + switch { + case depth < 0: + return ErrUnexpectedToken + case depth == 0: + return nil + } + } +} + +func open(dec *json.Decoder, want json.Delim) error { + tok, err := dec.Token() + if err != nil { + return err + } + if tok == nil { + return ErrNull + } + if d, ok := tok.(json.Delim); !ok || d != want { + return ErrUnexpectedToken + } + return nil +} + +func closeDelim(dec *json.Decoder, want json.Delim) error { + tok, err := dec.Token() + if err != nil { + return err + } + if d, ok := tok.(json.Delim); !ok || d != want { + return ErrUnexpectedToken + } + return nil +} + +// JSONTagString is a BuilderOptions.TagString for JSON input: it unquotes JSON +// strings and keeps other literals, such as numbers, as their raw text. +func JSONTagString(_ timeseries.FieldDefinition, raw []byte) string { + if isQuoted(raw) { + if b, err := unquote(raw); err == nil { + return string(b) + } + } + return string(raw) +} diff --git a/pkg/timeseries/dataset/stream/json_test.go b/pkg/timeseries/dataset/stream/json_test.go new file mode 100644 index 000000000..4d4bcbe48 --- /dev/null +++ b/pkg/timeseries/dataset/stream/json_test.go @@ -0,0 +1,298 @@ +/* + * Copyright 2018 The Trickster Authors + * + * 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 stream + +import ( + "encoding/json" + "io" + "strings" + "testing" + "testing/iotest" + + "github.com/trickstercache/trickster/v2/pkg/timeseries" + + "github.com/stretchr/testify/require" +) + +func sumWalk(sum *int64) func(*json.Decoder) error { + // sums the numbers in {"n":[...]}, skipping any other keys + return func(dec *json.Decoder) error { + return Object(dec, func(key string) error { + if key != "n" { + return Skip(dec) + } + return Array(dec, func() error { + var n json.Number + if err := dec.Decode(&n); err != nil { + return err + } + v, err := n.Int64() + *sum += v + return err + }) + }) + } +} + +func TestJSONFeeds(t *testing.T) { + const body = ` {"x":{"y":[1,{"z":null}]},"n":[1,2,3],"s":"str"} ` + var sum int64 + j := NewJSON(sumWalk(&sum), finishEmpty) + n, err := j.Write([]byte(body[:5])) + require.NoError(t, err) + require.Equal(t, 5, n) + _, err = j.Write([]byte(body[5:])) + require.NoError(t, err) + require.Zero(t, sum) + _, err = j.Finish() + require.NoError(t, err) + require.Equal(t, int64(6), sum) + + sum = 0 + j = NewJSON(sumWalk(&sum), finishEmpty) + _, err = j.Write([]byte(body[:10])) + require.NoError(t, err) + read, err := j.ReadFrom(iotest.OneByteReader(strings.NewReader(body[10:]))) + require.NoError(t, err) + require.Equal(t, int64(len(body)-10), read) + require.Equal(t, int64(6), sum) + _, err = j.Write([]byte(" ")) + require.ErrorIs(t, err, ErrInputConsumed) + _, err = j.ReadFrom(strings.NewReader(" ")) + require.ErrorIs(t, err, ErrInputConsumed) + _, err = j.Finish() + require.NoError(t, err) + _, err = j.Finish() + require.ErrorIs(t, err, ErrFinished) + _, err = j.Write(nil) + require.ErrorIs(t, err, ErrFinished) +} + +func TestJSONErrors(t *testing.T) { + tests := []struct { + name string + body string + err error + }{ + {"empty", "", io.ErrUnexpectedEOF}, + {"truncated", `{"n":[1,`, nil}, + {"truncated value", `{"n":[1`, nil}, + {"trailing", `{"n":[1]} 2`, ErrTrailingData}, + {"null object", `null`, ErrNull}, + {"null array", `{"n":null}`, ErrNull}, + {"wrong object", `[1]`, ErrUnexpectedToken}, + {"wrong array", `{"n":{}}`, ErrUnexpectedToken}, + {"number", `{"n":[1.5]}`, nil}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + var sum int64 + j := NewJSON(sumWalk(&sum), finishEmpty) + _, err := j.ReadFrom(strings.NewReader(test.body)) + requireError(t, err, test.err) + _, err2 := j.Finish() + require.Equal(t, err, err2) + + j = NewJSON(sumWalk(&sum), finishEmpty) + _, err = j.Write([]byte(test.body)) + require.NoError(t, err) + _, err = j.Finish() + requireError(t, err, test.err) + _, err = j.Write(nil) + require.ErrorIs(t, err, ErrFinished) + }) + } + require.ErrorIs(t, ErrNull, timeseries.ErrInvalidBody) + require.ErrorIs(t, ErrTrailingData, timeseries.ErrInvalidBody) +} + +func requireError(t *testing.T, err, want error) { + t.Helper() + if want == nil { + require.Error(t, err) + return + } + require.ErrorIs(t, err, want) +} + +func TestJSONStickyError(t *testing.T) { + var sum int64 + j := NewJSON(sumWalk(&sum), finishEmpty) + _, err := j.ReadFrom(strings.NewReader(`[`)) + require.ErrorIs(t, err, ErrUnexpectedToken) + _, err = j.Write(nil) + require.ErrorIs(t, err, ErrUnexpectedToken) + _, err = j.ReadFrom(strings.NewReader(`{}`)) + require.ErrorIs(t, err, ErrUnexpectedToken) +} + +func TestJSONVisitorMisuse(t *testing.T) { + walks := map[string]func(*json.Decoder) error{ + "object value not consumed": func(dec *json.Decoder) error { + return Object(dec, func(string) error { return nil }) + }, + "array element not consumed": func(dec *json.Decoder) error { + return Object(dec, func(string) error { + return Array(dec, func() error { return nil }) + }) + }, + } + for name, walk := range walks { + t.Run(name, func(t *testing.T) { + j := NewJSON(walk, finishEmpty) + _, err := j.ReadFrom(strings.NewReader(`{"a":[1]}`)) + require.ErrorIs(t, err, ErrValueNotConsumed) + }) + } + // a visitor that consumes only part of a value leaves the walk misaligned + partial := func(dec *json.Decoder) error { + return Object(dec, func(string) error { + return Array(dec, func() error { + _, err := dec.Token() + return err + }) + }) + } + for _, body := range []string{`{"a":[{"b":1}]}`, `{"a":[[1,2]]}`} { + j := NewJSON(partial, finishEmpty) + _, err := j.ReadFrom(strings.NewReader(body)) + require.ErrorIs(t, err, ErrUnexpectedToken, body) + } + keyed := func(dec *json.Decoder) error { + return Object(dec, func(string) error { + _, err := dec.Token() + return err + }) + } + j := NewJSON(keyed, finishEmpty) + _, err := j.ReadFrom(strings.NewReader(`{"a":[1,2]}`)) + require.ErrorIs(t, err, ErrUnexpectedToken) +} + +func TestJSONNestedErrors(t *testing.T) { + walk := func(dec *json.Decoder) error { + return Object(dec, func(string) error { + return Array(dec, func() error { return Skip(dec) }) + }) + } + for _, body := range []string{`{"a":[1,}`, `{"a":[1}`, `{"a":[1] "b"}`, `{`} { + j := NewJSON(walk, finishEmpty) + _, err := j.ReadFrom(strings.NewReader(body)) + require.Error(t, err, body) + } +} + +func TestJSONTagString(t *testing.T) { + var fd timeseries.FieldDefinition + require.Equal(t, "a", JSONTagString(fd, []byte(`"a"`))) + require.Equal(t, `a"b`, JSONTagString(fd, []byte(`"a\"b"`))) + require.Equal(t, "é", JSONTagString(fd, []byte(`"é"`))) + require.Equal(t, "12.5", JSONTagString(fd, []byte(`12.5`))) + require.Equal(t, "null", JSONTagString(fd, []byte(`null`))) + require.Equal(t, `"\x"`, JSONTagString(fd, []byte(`"\x"`))) + require.Equal(t, `"`, JSONTagString(fd, []byte(`"`))) +} + +type windowReader struct { + r io.Reader + dec *json.Decoder + read int64 + peak int64 +} + +func (w *windowReader) Read(p []byte) (int, error) { + // bytes read but not yet consumed are what the decoder is holding + if w.dec != nil { + w.peak = max(w.peak, w.read-w.dec.InputOffset()) + } + n, err := w.r.Read(p[:min(len(p), 512)]) + w.read += int64(n) + return n, err +} + +func skipPeak(t *testing.T, body string, skip func(*json.Decoder) error) int64 { + t.Helper() + w := &windowReader{r: strings.NewReader(body)} + var sum int64 + walk := func(dec *json.Decoder) error { + w.dec = dec + return Object(dec, func(key string) error { + if key != "n" { + return skip(dec) + } + return Array(dec, func() error { + var n int64 + err := dec.Decode(&n) + sum += n + return err + }) + }) + } + _, err := NewJSON(walk, finishEmpty).ReadFrom(w) + require.NoError(t, err) + require.Equal(t, int64(3), sum) + return w.peak +} + +func TestSkipDoesNotHoldValue(t *testing.T) { + var sb strings.Builder + sb.WriteString(`{"skip":[`) + for i := range 50000 { + if i > 0 { + sb.WriteByte(',') + } + sb.WriteString(`{"k":"value","n":[1,2,3]}`) + } + sb.WriteString(`],"n":[1,2]}`) + body := sb.String() + require.Less(t, skipPeak(t, body, Skip), int64(64<<10)) + // decoding the value whole holds nearly all of it, which shows the measurement works + whole := func(dec *json.Decoder) error { return dec.Decode(new(json.RawMessage)) } + require.Greater(t, skipPeak(t, body, whole), int64(len(body)/2)) +} + +func TestSkip(t *testing.T) { + tests := []struct { + body string + err error + }{ + {`"str"`, nil}, + {`12.5`, nil}, + {`null`, nil}, + {`{"a":{"b":[1,{"c":null}]},"d":"e"}`, nil}, + {`[[],{},[[{}]]]`, nil}, + {`[1,[2`, io.ErrUnexpectedEOF}, + } + for _, test := range tests { + j := NewJSON(Skip, finishEmpty) + _, err := j.ReadFrom(strings.NewReader(test.body)) + if test.err == nil { + require.NoError(t, err, test.body) + continue + } + require.Error(t, err, test.body) + } + // with no value left to skip, Skip meets the closing delimiter instead + closing := func(dec *json.Decoder) error { + if _, err := dec.Token(); err != nil { + return err + } + return Skip(dec) + } + _, err := NewJSON(closing, finishEmpty).ReadFrom(strings.NewReader(`[]`)) + require.ErrorIs(t, err, ErrUnexpectedToken) +} diff --git a/pkg/timeseries/dataset/stream/lines.go b/pkg/timeseries/dataset/stream/lines.go new file mode 100644 index 000000000..74d04a78b --- /dev/null +++ b/pkg/timeseries/dataset/stream/lines.go @@ -0,0 +1,216 @@ +/* + * Copyright 2018 The Trickster Authors + * + * 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 stream + +import ( + "bytes" + "errors" + "fmt" + "io" + "sync" + + "github.com/trickstercache/trickster/v2/pkg/timeseries" +) + +// DefaultMaxLineBytes is the longest line, excluding its terminator, that a +// Lines decoder accepts unless SetMaxLineBytes is used. +const DefaultMaxLineBytes = 16 << 20 + +const readBufferSize = 32 << 10 + +// ErrLineTooLong indicates a line exceeded the decoder's maximum line length. +var ErrLineTooLong = fmt.Errorf("%w: line exceeds maximum length", timeseries.ErrInvalidBody) + +var readBuffers = sync.Pool{New: func() any { + b := make([]byte, readBufferSize) + return &b +}} + +// Lines is a Decoder for newline-delimited formats such as TSV and JSON Lines. +// Each line is passed to its callback without the "\n" or "\r\n" terminator. +type Lines struct { + onLine func(line []byte) error + finish FinishFunc + partial []byte + max int + err error + done bool +} + +var _ Decoder = (*Lines)(nil) + +// NewLines returns a Lines decoder. onLine must not retain line after returning, +// and finish is called by Finish after the last line. +func NewLines(onLine func(line []byte) error, finish FinishFunc) *Lines { + return &Lines{onLine: onLine, finish: finish, max: DefaultMaxLineBytes} +} + +// SetMaxLineBytes sets the longest line, excluding its terminator, that the +// decoder accepts. Values below 1 restore DefaultMaxLineBytes. +func (l *Lines) SetMaxLineBytes(n int) *Lines { + l.max = n + if n < 1 { + l.max = DefaultMaxLineBytes + } + return l +} + +// Write passes each complete line in p to the callback and holds any trailing +// partial line until more input or Finish arrives. +func (l *Lines) Write(p []byte) (int, error) { + if err := l.check(); err != nil { + return 0, err + } + rest := p + for len(rest) > 0 { + i := bytes.IndexByte(rest, '\n') + if i < 0 { + if !l.fits(rest) { + return len(p) - len(rest), l.fail(ErrLineTooLong) + } + l.partial = append(l.partial, rest...) + break + } + line := rest[:i] + if len(l.partial) > 0 { + // check before joining, so an overlong line never grows the partial buffer + if !l.fits(line) { + return len(p) - len(rest), l.fail(ErrLineTooLong) + } + l.partial = append(l.partial, line...) + line = l.partial + } + if err := l.emit(line); err != nil { + return len(p) - len(rest), err + } + l.partial = l.partial[:0] + rest = rest[i+1:] + } + return len(p), nil +} + +// ReadFrom reads r to EOF, passing each line to the callback. A reader that +// implements io.WriterTo writes into the decoder directly, avoiding a copy. +func (l *Lines) ReadFrom(r io.Reader) (int64, error) { + if err := l.check(); err != nil { + return 0, err + } + if wt, ok := r.(io.WriterTo); ok { + n, err := wt.WriteTo(l) + if err != nil { + return n, l.fail(err) + } + return n, nil + } + bp := readBuffers.Get().(*[]byte) + defer readBuffers.Put(bp) + buf := *bp + var total int64 + for { + n, err := r.Read(buf) + if n > 0 { + total += int64(n) + if _, werr := l.Write(buf[:n]); werr != nil { + return total, werr + } + } + if errors.Is(err, io.EOF) { + return total, nil + } + if err != nil { + return total, l.fail(err) + } + } +} + +// Finish passes any final unterminated line to the callback, then calls the +// finish function. +func (l *Lines) Finish() (timeseries.Timeseries, error) { + if l.done { + return nil, ErrFinished + } + l.done = true + if l.err != nil { + return nil, l.err + } + if len(l.partial) > 0 { + if err := l.emit(l.partial); err != nil { + return nil, err + } + } + l.partial = nil + return l.finish() +} + +func (l *Lines) check() error { + if l.done { + return ErrFinished + } + return l.err +} + +func (l *Lines) fail(err error) error { + if l.err == nil { + l.err = err + } + return l.err +} + +func (l *Lines) fits(next []byte) bool { + // the partial line plus next, less a trailing "\r" whose "\n" may follow + n := len(l.partial) + len(next) + var last byte + switch { + case len(next) > 0: + last = next[len(next)-1] + case len(l.partial) > 0: + last = l.partial[len(l.partial)-1] + } + if last == '\r' { + n-- + } + return n <= l.max +} + +func (l *Lines) emit(raw []byte) error { + line := raw + if n := len(line); n > 0 && line[n-1] == '\r' { + line = line[:n-1] + } + if len(line) > l.max { + return l.fail(ErrLineTooLong) + } + if err := l.onLine(line); err != nil { + return l.fail(err) + } + return nil +} + +// SplitFields splits line at each sep into dst, reusing its capacity. Quotes +// and escapes are not interpreted, and the fields alias line. +func SplitFields(line []byte, sep byte, dst [][]byte) [][]byte { + out := dst[:0] + rest := line + for { + i := bytes.IndexByte(rest, sep) + if i < 0 { + return append(out, rest) + } + out = append(out, rest[:i]) + rest = rest[i+1:] + } +} diff --git a/pkg/timeseries/dataset/stream/lines_test.go b/pkg/timeseries/dataset/stream/lines_test.go new file mode 100644 index 000000000..ebf34871a --- /dev/null +++ b/pkg/timeseries/dataset/stream/lines_test.go @@ -0,0 +1,228 @@ +/* + * Copyright 2018 The Trickster Authors + * + * 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 stream + +import ( + "bytes" + "errors" + "io" + "strings" + "testing" + "testing/iotest" + + "github.com/trickstercache/trickster/v2/pkg/timeseries" + "github.com/trickstercache/trickster/v2/pkg/timeseries/dataset" + + "github.com/stretchr/testify/require" +) + +var errCallback = errors.New("callback error") + +func collectLines(lines *[]string) func([]byte) error { + return func(line []byte) error { + *lines = append(*lines, string(line)) + return nil + } +} + +func finishEmpty() (timeseries.Timeseries, error) { + return &dataset.DataSet{}, nil +} + +func TestLinesWrite(t *testing.T) { + tests := []struct { + name string + chunks []string + want []string + }{ + {"single", []string{"a\nb\n"}, []string{"a", "b"}}, + {"unterminated", []string{"a\nb"}, []string{"a", "b"}}, + {"crlf", []string{"a\r\nb\r\n"}, []string{"a", "b"}}, + {"split crlf", []string{"a\r", "\nb"}, []string{"a", "b"}}, + {"split lines", []string{"ab", "c\nd", "e", "\n", "f"}, []string{"abc", "de", "f"}}, + {"empty lines", []string{"\n\na\n", "\n"}, []string{"", "", "a", ""}}, + {"final cr", []string{"a\r"}, []string{"a"}}, + {"nothing", nil, nil}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + var got []string + l := NewLines(collectLines(&got), finishEmpty) + for _, c := range test.chunks { + n, err := l.Write([]byte(c)) + require.NoError(t, err) + require.Equal(t, len(c), n) + } + ts, err := l.Finish() + require.NoError(t, err) + require.NotNil(t, ts) + require.Equal(t, test.want, got) + }) + } +} + +func TestLinesMaxLineBytes(t *testing.T) { + var got []string + l := NewLines(collectLines(&got), finishEmpty).SetMaxLineBytes(3) + _, err := l.Write([]byte("abc\r")) + require.NoError(t, err) + _, err = l.Write([]byte("\nab")) + require.NoError(t, err) + // only a trailing "\r" may take a partial line past the limit + n, err := l.Write([]byte("cd")) + require.ErrorIs(t, err, ErrLineTooLong) + require.ErrorIs(t, err, timeseries.ErrInvalidBody) + require.Zero(t, n) + _, err = l.Write([]byte("x")) + require.ErrorIs(t, err, ErrLineTooLong) + _, err = l.Finish() + require.ErrorIs(t, err, ErrLineTooLong) + require.Equal(t, []string{"abc"}, got) + + l = NewLines(collectLines(&got), finishEmpty).SetMaxLineBytes(3) + n, err = l.Write([]byte("ab\nabcd\nab\n")) + require.ErrorIs(t, err, ErrLineTooLong) + require.Equal(t, 3, n) + + require.Equal(t, DefaultMaxLineBytes, NewLines(nil, nil).SetMaxLineBytes(0).max) +} + +func TestLinesLimitPrecedesGrowth(t *testing.T) { + var got []string + l := NewLines(collectLines(&got), finishEmpty).SetMaxLineBytes(8) + _, err := l.Write([]byte("abc")) + require.NoError(t, err) + huge := append(bytes.Repeat([]byte{'x'}, 1<<20), '\n') + n, err := l.Write(huge) + require.ErrorIs(t, err, ErrLineTooLong) + require.Zero(t, n) + require.Less(t, cap(l.partial), 1<<10) + + // a "\r" held from one write still ends the line when its "\n" arrives + l = NewLines(collectLines(&got), finishEmpty).SetMaxLineBytes(3) + for _, chunk := range []string{"ab", "c\r", "\n"} { + _, err = l.Write([]byte(chunk)) + require.NoError(t, err) + } + _, err = l.Write([]byte("ab")) + require.NoError(t, err) + _, err = l.Write([]byte("cd\n")) + require.ErrorIs(t, err, ErrLineTooLong) + require.Equal(t, []string{"abc"}, got) +} + +func TestLinesCallbackError(t *testing.T) { + calls := 0 + l := NewLines(func([]byte) error { + calls++ + return errCallback + }, finishEmpty) + n, err := l.Write([]byte("a\nb\n")) + require.ErrorIs(t, err, errCallback) + require.Zero(t, n) + _, err = l.Write([]byte("c\n")) + require.ErrorIs(t, err, errCallback) + _, err = l.ReadFrom(strings.NewReader("d\n")) + require.ErrorIs(t, err, errCallback) + require.Equal(t, 1, calls) + + // a failing final line surfaces from Finish + l = NewLines(func([]byte) error { return errCallback }, finishEmpty) + _, err = l.Write([]byte("a")) + require.NoError(t, err) + _, err = l.Finish() + require.ErrorIs(t, err, errCallback) +} + +func TestLinesFinished(t *testing.T) { + l := NewLines(func([]byte) error { return nil }, finishEmpty) + _, err := l.Finish() + require.NoError(t, err) + _, err = l.Finish() + require.ErrorIs(t, err, ErrFinished) + _, err = l.Write([]byte("a")) + require.ErrorIs(t, err, ErrFinished) + _, err = l.ReadFrom(strings.NewReader("a")) + require.ErrorIs(t, err, ErrFinished) +} + +func TestLinesReadFrom(t *testing.T) { + body := "a\nbb\r\nccc" + readers := map[string]io.Reader{ + "writer-to": bytes.NewReader([]byte(body)), + "reader": iotest.HalfReader(strings.NewReader(body)), + "data-err": iotest.DataErrReader(strings.NewReader(body)), + } + for name, r := range readers { + t.Run(name, func(t *testing.T) { + var got []string + l := NewLines(collectLines(&got), finishEmpty) + n, err := l.ReadFrom(r) + require.NoError(t, err) + require.Equal(t, int64(len(body)), n) + _, err = l.Finish() + require.NoError(t, err) + require.Equal(t, []string{"a", "bb", "ccc"}, got) + }) + } +} + +type errWriterTo struct{} + +func (errWriterTo) Read([]byte) (int, error) { return 0, io.EOF } +func (errWriterTo) WriteTo(io.Writer) (int64, error) { return 0, io.ErrShortWrite } + +func TestLinesReadFromErrors(t *testing.T) { + l := NewLines(func([]byte) error { return nil }, finishEmpty) + _, err := l.ReadFrom(io.MultiReader(strings.NewReader("a\n"), iotest.ErrReader(errCallback))) + require.ErrorIs(t, err, errCallback) + _, err = l.Finish() + require.ErrorIs(t, err, errCallback) + + l = NewLines(func([]byte) error { return nil }, finishEmpty) + _, err = l.ReadFrom(errWriterTo{}) + require.ErrorIs(t, err, io.ErrShortWrite) + + l = NewLines(func([]byte) error { return errCallback }, finishEmpty) + _, err = l.ReadFrom(iotest.OneByteReader(strings.NewReader("a\nb\n"))) + require.ErrorIs(t, err, errCallback) +} + +func TestSplitFields(t *testing.T) { + tests := []struct { + line string + want []string + }{ + {"", []string{""}}, + {"a", []string{"a"}}, + {"a\tb\t\tc", []string{"a", "b", "", "c"}}, + {"a\t", []string{"a", ""}}, + } + var dst [][]byte + for _, test := range tests { + dst = SplitFields([]byte(test.line), '\t', dst) + got := make([]string, len(dst)) + for i, f := range dst { + got[i] = string(f) + } + require.Equal(t, test.want, got) + } + line := []byte("a\tb\tc") + dst = make([][]byte, 0, 4) + allocs := testing.AllocsPerRun(10, func() { dst = SplitFields(line, '\t', dst) }) + require.Zero(t, allocs) +} diff --git a/pkg/timeseries/dataset/stream/stream.go b/pkg/timeseries/dataset/stream/stream.go new file mode 100644 index 000000000..569d93e86 --- /dev/null +++ b/pkg/timeseries/dataset/stream/stream.go @@ -0,0 +1,87 @@ +/* + * Copyright 2018 The Trickster Authors + * + * 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 stream decodes upstream response bodies into Timeseries in a single +// pass, without first unmarshaling the whole body into an intermediate model. +package stream + +import ( + "bytes" + "errors" + "io" + + "github.com/trickstercache/trickster/v2/pkg/timeseries" +) + +// Decoder consumes an upstream response body and builds a Timeseries from it. +// Feed it with Write calls and/or ReadFrom, then call Finish once. +type Decoder interface { + io.Writer + io.ReaderFrom + // Finish completes the decode and returns the resulting Timeseries. + Finish() (timeseries.Timeseries, error) +} + +// NewDecoderFunc returns a Decoder for one response to the provided query. +type NewDecoderFunc func(*timeseries.TimeRangeQuery) (Decoder, error) + +// FinishFunc completes a decode once all input has been consumed. +type FinishFunc func() (timeseries.Timeseries, error) + +var ( + // ErrFinished indicates a Decoder was used after Finish. + ErrFinished = errors.New("decoder already finished") + // ErrInputConsumed indicates input was provided after a Decoder read its input to the end. + ErrInputConsumed = errors.New("decoder input already consumed") +) + +// ReaderUnmarshaler adapts newDecoder to a timeseries.UnmarshalerReaderFunc. It +// feeds the reader to the Decoder's ReadFrom, so the body is decoded as it is read. +func ReaderUnmarshaler(newDecoder NewDecoderFunc) timeseries.UnmarshalerReaderFunc { + return func(r io.Reader, trq *timeseries.TimeRangeQuery) (timeseries.Timeseries, error) { + if r == nil { + return nil, timeseries.ErrInvalidBody + } + return decode(newDecoder, r, trq) + } +} + +// BytesUnmarshaler adapts newDecoder to a timeseries.UnmarshalerFunc. +func BytesUnmarshaler(newDecoder NewDecoderFunc) timeseries.UnmarshalerFunc { + return func(b []byte, trq *timeseries.TimeRangeQuery) (timeseries.Timeseries, error) { + return decode(newDecoder, bytes.NewReader(b), trq) + } +} + +func decode(newDecoder NewDecoderFunc, r io.Reader, + trq *timeseries.TimeRangeQuery, +) (timeseries.Timeseries, error) { + dec, err := newDecoder(trq) + if err != nil { + return nil, err + } + if _, err = dec.ReadFrom(r); err != nil { + return nil, err + } + ts, err := dec.Finish() + if err != nil { + return nil, err + } + if ts == nil { + return nil, timeseries.ErrInvalidBody + } + return ts, nil +} diff --git a/pkg/timeseries/dataset/stream/streamtest/compare.go b/pkg/timeseries/dataset/stream/streamtest/compare.go new file mode 100644 index 000000000..512c2b979 --- /dev/null +++ b/pkg/timeseries/dataset/stream/streamtest/compare.go @@ -0,0 +1,194 @@ +/* + * Copyright 2018 The Trickster Authors + * + * 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 streamtest + +import ( + "bytes" + "fmt" + "maps" + "math" + "reflect" + "slices" + "strings" + + "github.com/trickstercache/trickster/v2/pkg/timeseries" + "github.com/trickstercache/trickster/v2/pkg/timeseries/dataset" +) + +// CompareOptions relaxes Compare. +type CompareOptions struct { + // IgnoreSizes skips the size fields of points, series and series headers. + IgnoreSizes bool + // IgnoreSeriesOrder matches each result's series by name and tags, not position. + IgnoreSeriesOrder bool +} + +// Compare returns an error describing the first difference between want and got, +// or nil when they match. Float NaN values are equal to each other. +func Compare(want, got *dataset.DataSet, o CompareOptions) error { + if want == nil || got == nil { + if want == nil && got == nil { + return nil + } + return fmt.Errorf("dataset: want nil %t, got nil %t", want == nil, got == nil) + } + for _, f := range []struct{ name, want, got string }{ + {"sourceResultType", want.SourceResultType, got.SourceResultType}, + {"status", want.Status, got.Status}, + {"error", want.Error, got.Error}, + {"errorType", want.ErrorType, got.ErrorType}, + } { + if f.want != f.got { + return fmt.Errorf("%s: want %q, got %q", f.name, f.want, f.got) + } + } + if !slices.Equal(want.Warnings, got.Warnings) { + return fmt.Errorf("warnings: want %q, got %q", want.Warnings, got.Warnings) + } + if err := compareExtents("extents", want.ExtentList, got.ExtentList); err != nil { + return err + } + if err := compareExtents("volatileExtents", want.VolatileExtentList, got.VolatileExtentList); err != nil { + return err + } + if len(want.Results) != len(got.Results) { + return fmt.Errorf("results: want %d, got %d", len(want.Results), len(got.Results)) + } + for i := range want.Results { + if err := compareResult(want.Results[i], got.Results[i], o); err != nil { + return fmt.Errorf("results[%d].%w", i, err) + } + } + return nil +} + +func compareExtents(name string, want, got timeseries.ExtentList) error { + if len(want) != len(got) { + return fmt.Errorf("%s: want %v, got %v", name, want, got) + } + for i := range want { + if !want[i].Start.Equal(got[i].Start) || !want[i].End.Equal(got[i].End) { + return fmt.Errorf("%s: want %v, got %v", name, want, got) + } + } + return nil +} + +func compareResult(want, got *dataset.Result, o CompareOptions) error { + if want == nil || got == nil { + if want == nil && got == nil { + return nil + } + return fmt.Errorf("result: want nil %t, got nil %t", want == nil, got == nil) + } + if want.StatementID != got.StatementID || want.Name != got.Name || want.Error != got.Error { + return fmt.Errorf("result: want {%d %q %q}, got {%d %q %q}", want.StatementID, want.Name, + want.Error, got.StatementID, got.Name, got.Error) + } + if len(want.SeriesList) != len(got.SeriesList) { + return fmt.Errorf("series: want %d, got %d", len(want.SeriesList), len(got.SeriesList)) + } + ws, gs := want.SeriesList, got.SeriesList + if o.IgnoreSeriesOrder { + ws, gs = sortedByKey(ws), sortedByKey(gs) + } + for i := range ws { + if err := compareSeries(ws[i], gs[i], o); err != nil { + return fmt.Errorf("series[%d]%w", i, err) + } + } + return nil +} + +func sortedByKey(sl dataset.SeriesList) dataset.SeriesList { + keys := make(map[*dataset.Series]string, len(sl)) + for _, s := range sl { + if s != nil { + keys[s] = s.Header.Name + "\x00" + s.Header.Tags.JSON() + } + } + out := slices.Clone(sl) + slices.SortStableFunc(out, func(a, b *dataset.Series) int { + return strings.Compare(keys[a], keys[b]) + }) + return out +} + +func compareSeries(want, got *dataset.Series, o CompareOptions) error { + if want == nil || got == nil { + if want == nil && got == nil { + return nil + } + return fmt.Errorf(": want nil %t, got nil %t", want == nil, got == nil) + } + wh, gh := &want.Header, &got.Header + switch { + case wh.Name != gh.Name: + return fmt.Errorf(".name: want %q, got %q", wh.Name, gh.Name) + case wh.QueryStatement != gh.QueryStatement: + return fmt.Errorf(".query: want %q, got %q", wh.QueryStatement, gh.QueryStatement) + case !maps.Equal(wh.Tags, gh.Tags): + return fmt.Errorf(".tags: want %v, got %v", wh.Tags, gh.Tags) + case wh.TimestampField != gh.TimestampField: + return fmt.Errorf(".timestampField: want %v, got %v", wh.TimestampField, gh.TimestampField) + case !slices.Equal(wh.TagFieldsList, gh.TagFieldsList): + return fmt.Errorf(".tagFields: want %v, got %v", wh.TagFieldsList, gh.TagFieldsList) + case !slices.Equal(wh.ValueFieldsList, gh.ValueFieldsList): + return fmt.Errorf(".valueFields: want %v, got %v", wh.ValueFieldsList, gh.ValueFieldsList) + case !slices.Equal(wh.UntrackedFieldsList, gh.UntrackedFieldsList): + return fmt.Errorf(".untrackedFields: want %v, got %v", wh.UntrackedFieldsList, gh.UntrackedFieldsList) + case !o.IgnoreSizes && wh.Size != gh.Size: + return fmt.Errorf(".headerSize: want %d, got %d", wh.Size, gh.Size) + case !o.IgnoreSizes && want.PointSize != got.PointSize: + return fmt.Errorf(".pointSize: want %d, got %d", want.PointSize, got.PointSize) + case len(want.Points) != len(got.Points): + return fmt.Errorf(".points: want %d, got %d", len(want.Points), len(got.Points)) + } + for i := range want.Points { + wp, gp := &want.Points[i], &got.Points[i] + switch { + case wp.Epoch != gp.Epoch: + return fmt.Errorf(".points[%d].epoch: want %d, got %d", i, wp.Epoch, gp.Epoch) + case !o.IgnoreSizes && wp.Size != gp.Size: + return fmt.Errorf(".points[%d].size: want %d, got %d", i, wp.Size, gp.Size) + case len(wp.Values) != len(gp.Values): + return fmt.Errorf(".points[%d].values: want %v, got %v", i, wp.Values, gp.Values) + } + for j := range wp.Values { + if !valueEqual(wp.Values[j], gp.Values[j]) { + return fmt.Errorf(".points[%d].values[%d]: want %#v, got %#v", i, j, + wp.Values[j], gp.Values[j]) + } + } + } + return nil +} + +func valueEqual(want, got any) bool { + switch w := want.(type) { + case float64: + g, ok := got.(float64) + return ok && (w == g || (math.IsNaN(w) && math.IsNaN(g))) + case float32: + g, ok := got.(float32) + return ok && (w == g || (math.IsNaN(float64(w)) && math.IsNaN(float64(g)))) + case []byte: + g, ok := got.([]byte) + return ok && bytes.Equal(w, g) + } + return reflect.DeepEqual(want, got) +} diff --git a/pkg/timeseries/dataset/stream/streamtest/compare_test.go b/pkg/timeseries/dataset/stream/streamtest/compare_test.go new file mode 100644 index 000000000..581692c1c --- /dev/null +++ b/pkg/timeseries/dataset/stream/streamtest/compare_test.go @@ -0,0 +1,142 @@ +/* + * Copyright 2018 The Trickster Authors + * + * 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 streamtest + +import ( + "math" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/timeseries" + "github.com/trickstercache/trickster/v2/pkg/timeseries/dataset" + + "github.com/stretchr/testify/require" +) + +func compareBase() *dataset.DataSet { + fd := timeseries.FieldDefinition{Name: "v", DataType: timeseries.Float64} + series := func(host string, values ...any) *dataset.Series { + pts := make(dataset.Points, len(values)) + for i, v := range values { + pts[i] = dataset.Point{Epoch: 1, Size: 10, Values: []any{v}} + } + return &dataset.Series{ + Header: dataset.SeriesHeader{Name: "s", Tags: dataset.Tags{"host": host}, + ValueFieldsList: timeseries.FieldDefinitions{fd}, Size: 5}, + Points: pts, + PointSize: 10, + } + } + return &dataset.DataSet{ + Status: "success", + Warnings: []string{"w"}, + ExtentList: timeseries.ExtentList{{Start: time.Unix(1, 0), End: time.Unix(2, 0)}}, + Results: dataset.Results{{ + StatementID: 1, + SeriesList: dataset.SeriesList{ + series("a", math.NaN(), float32(math.NaN()), []byte("b"), map[string]int{"k": 1}), + series("b", 1.0), + nil, + }, + }, nil}, + } +} + +func TestCompare(t *testing.T) { + require.NoError(t, Compare(nil, nil, CompareOptions{})) + require.Error(t, Compare(compareBase(), nil, CompareOptions{})) + require.NoError(t, Compare(compareBase(), compareBase(), CompareOptions{})) + + tests := []struct { + name string + mutate func(*dataset.DataSet) + want string + opts CompareOptions + }{ + {"status", func(ds *dataset.DataSet) { ds.Status = "error" }, "status:", CompareOptions{}}, + {"source", func(ds *dataset.DataSet) { ds.SourceResultType = "matrix" }, "sourceResultType:", CompareOptions{}}, + {"warnings", func(ds *dataset.DataSet) { ds.Warnings = nil }, "warnings:", CompareOptions{}}, + {"extent count", func(ds *dataset.DataSet) { ds.ExtentList = nil }, "extents:", CompareOptions{}}, + {"extent", func(ds *dataset.DataSet) { ds.ExtentList[0].End = time.Unix(3, 0) }, "extents:", CompareOptions{}}, + {"volatile", func(ds *dataset.DataSet) { ds.VolatileExtentList = ds.ExtentList }, "volatileExtents:", CompareOptions{}}, + {"results", func(ds *dataset.DataSet) { ds.Results = ds.Results[:1] }, "results: want 2, got 1", CompareOptions{}}, + {"nil result", func(ds *dataset.DataSet) { ds.Results[0] = nil }, "results[0].result: want nil false", CompareOptions{}}, + {"statement", func(ds *dataset.DataSet) { ds.Results[0].StatementID = 2 }, "results[0].result: want {1", CompareOptions{}}, + {"series count", func(ds *dataset.DataSet) { ds.Results[0].SeriesList = ds.Results[0].SeriesList[:1] }, + "results[0].series: want 3, got 1", CompareOptions{}}, + {"nil series", func(ds *dataset.DataSet) { ds.Results[0].SeriesList[1] = nil }, "series[1]: want nil false", CompareOptions{}}, + {"order", func(ds *dataset.DataSet) { + sl := ds.Results[0].SeriesList + sl[0], sl[1] = sl[1], sl[0] + }, "series[0].tags", CompareOptions{}}, + {"name", func(ds *dataset.DataSet) { ds.Results[0].SeriesList[0].Header.Name = "x" }, ".name:", CompareOptions{}}, + {"query", func(ds *dataset.DataSet) { ds.Results[0].SeriesList[0].Header.QueryStatement = "x" }, ".query:", CompareOptions{}}, + {"timestamp", func(ds *dataset.DataSet) { ds.Results[0].SeriesList[0].Header.TimestampField.Name = "t" }, + ".timestampField:", CompareOptions{}}, + {"tag fields", func(ds *dataset.DataSet) { + ds.Results[0].SeriesList[0].Header.TagFieldsList = timeseries.FieldDefinitions{{Name: "host"}} + }, ".tagFields:", CompareOptions{}}, + {"value fields", func(ds *dataset.DataSet) { ds.Results[0].SeriesList[0].Header.ValueFieldsList = nil }, + ".valueFields:", CompareOptions{}}, + {"untracked fields", func(ds *dataset.DataSet) { + ds.Results[0].SeriesList[0].Header.UntrackedFieldsList = timeseries.FieldDefinitions{{Name: "u"}} + }, ".untrackedFields:", CompareOptions{}}, + {"header size", func(ds *dataset.DataSet) { ds.Results[0].SeriesList[0].Header.Size = 6 }, ".headerSize:", CompareOptions{}}, + {"point size", func(ds *dataset.DataSet) { ds.Results[0].SeriesList[0].PointSize = 6 }, ".pointSize:", CompareOptions{}}, + {"points", func(ds *dataset.DataSet) { ds.Results[0].SeriesList[1].Points = nil }, "series[1].points: want 1", CompareOptions{}}, + {"epoch", func(ds *dataset.DataSet) { ds.Results[0].SeriesList[1].Points[0].Epoch = 2 }, ".points[0].epoch:", CompareOptions{}}, + {"size", func(ds *dataset.DataSet) { ds.Results[0].SeriesList[1].Points[0].Size = 2 }, ".points[0].size:", CompareOptions{}}, + {"value count", func(ds *dataset.DataSet) { ds.Results[0].SeriesList[1].Points[0].Values = nil }, + ".points[0].values:", CompareOptions{}}, + {"float", func(ds *dataset.DataSet) { ds.Results[0].SeriesList[1].Points[0].Values[0] = 2.0 }, + ".points[0].values[0]: want 1, got 2", CompareOptions{}}, + {"float type", func(ds *dataset.DataSet) { ds.Results[0].SeriesList[1].Points[0].Values[0] = int64(1) }, + ".points[0].values[0]:", CompareOptions{}}, + {"nan", func(ds *dataset.DataSet) { ds.Results[0].SeriesList[0].Points[0].Values[0] = 1.0 }, + ".points[0].values[0]:", CompareOptions{}}, + {"float32", func(ds *dataset.DataSet) { ds.Results[0].SeriesList[0].Points[1].Values[0] = float32(1) }, + ".points[1].values[0]:", CompareOptions{}}, + {"bytes", func(ds *dataset.DataSet) { ds.Results[0].SeriesList[0].Points[2].Values[0] = []byte("c") }, + ".points[2].values[0]:", CompareOptions{}}, + {"other", func(ds *dataset.DataSet) { ds.Results[0].SeriesList[0].Points[3].Values[0] = map[string]int{} }, + ".points[3].values[0]:", CompareOptions{}}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + got := compareBase() + test.mutate(got) + err := Compare(compareBase(), got, test.opts) + require.ErrorContains(t, err, test.want) + }) + } +} + +func TestCompareOptions(t *testing.T) { + got := compareBase() + s := got.Results[0].SeriesList[0] + s.Header.Size, s.PointSize, s.Points[0].Size = 1, 2, 3 + require.Error(t, Compare(compareBase(), got, CompareOptions{})) + require.NoError(t, Compare(compareBase(), got, CompareOptions{IgnoreSizes: true})) + + got = compareBase() + sl := got.Results[0].SeriesList + sl[0], sl[1], sl[2] = sl[2], sl[0], sl[1] + require.Error(t, Compare(compareBase(), got, CompareOptions{})) + require.NoError(t, Compare(compareBase(), got, CompareOptions{IgnoreSeriesOrder: true})) + // reordering must not mutate either DataSet + require.Nil(t, got.Results[0].SeriesList[0]) +} diff --git a/pkg/timeseries/dataset/stream/streamtest/streamtest.go b/pkg/timeseries/dataset/stream/streamtest/streamtest.go new file mode 100644 index 000000000..522f836a7 --- /dev/null +++ b/pkg/timeseries/dataset/stream/streamtest/streamtest.go @@ -0,0 +1,285 @@ +/* + * Copyright 2018 The Trickster Authors + * + * 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 streamtest checks that stream.Decoder implementations produce the same +// result however their input arrives, and benchmarks them against other unmarshalers. +package streamtest + +import ( + "bytes" + "errors" + "io" + "slices" + "strconv" + "testing" + "testing/iotest" + + "github.com/trickstercache/trickster/v2/pkg/timeseries" + "github.com/trickstercache/trickster/v2/pkg/timeseries/dataset" + "github.com/trickstercache/trickster/v2/pkg/timeseries/dataset/stream" + "github.com/trickstercache/trickster/v2/pkg/util/weak/weaktest" +) + +// ErrAny, as a Case's WantErr, accepts any non-nil error. +var ErrAny = errors.New("any error") + +var errInjected = errors.New("injected read error") + +const shuffleRounds = 4 + +// Case describes one upstream response body for Conformance. +type Case struct { + // TRQ is passed to every decoder and unmarshaler. + TRQ *timeseries.TimeRangeQuery + // Body is the upstream response body. + Body []byte + // WantErr, when set, is the error every decode must return, per errors.Is. + WantErr error + // Want, when set, is the DataSet every decode must produce, ignoring sizes. + Want *dataset.DataSet + // Legacy, when set, is an existing unmarshaler whose DataSet must match, ignoring sizes. + Legacy timeseries.UnmarshalerReaderFunc + // Shuffle, when set, reorders Body without changing its meaning; see ShuffleLines. + Shuffle func(body []byte, rng *weaktest.Rand) []byte + // Unwrap returns the DataSet within a decoded Timeseries. When nil, the + // Timeseries must be a *dataset.DataSet. + Unwrap func(timeseries.Timeseries) *dataset.DataSet +} + +type outcome struct { + name string + ts timeseries.Timeseries + err error +} + +type feed struct { + name string + run func(dec stream.Decoder, body []byte) error +} + +type readerOnly struct{ io.Reader } + +var feeds = []feed{ + {"write-once", func(dec stream.Decoder, body []byte) error { + _, err := dec.Write(body) + return err + }}, + {"write-bytewise", func(dec stream.Decoder, body []byte) error { + for i := range body { + if _, err := dec.Write(body[i : i+1]); err != nil { + return err + } + } + return nil + }}, + {"write-random", func(dec stream.Decoder, body []byte) error { + rng := weaktest.NewRand(0x5eed, 0x5eed) + for rest := body; len(rest) > 0; { + n := min(1+rng.IntN(64), len(rest)) + if _, err := dec.Write(rest[:n]); err != nil { + return err + } + rest = rest[n:] + } + return nil + }}, + {"write-then-readfrom", func(dec stream.Decoder, body []byte) error { + half := len(body) / 2 + if _, err := dec.Write(body[:half]); err != nil { + return err + } + _, err := dec.ReadFrom(readerOnly{bytes.NewReader(body[half:])}) + return err + }}, + {"readfrom", func(dec stream.Decoder, body []byte) error { + _, err := dec.ReadFrom(bytes.NewReader(body)) + return err + }}, + {"readfrom-reader", func(dec stream.Decoder, body []byte) error { + _, err := dec.ReadFrom(readerOnly{bytes.NewReader(body)}) + return err + }}, + {"readfrom-onebyte", func(dec stream.Decoder, body []byte) error { + _, err := dec.ReadFrom(iotest.OneByteReader(bytes.NewReader(body))) + return err + }}, + {"readfrom-half", func(dec stream.Decoder, body []byte) error { + _, err := dec.ReadFrom(iotest.HalfReader(bytes.NewReader(body))) + return err + }}, + {"readfrom-dataerr", func(dec stream.Decoder, body []byte) error { + _, err := dec.ReadFrom(iotest.DataErrReader(bytes.NewReader(body))) + return err + }}, +} + +// Conformance decodes c.Body through every way a Decoder can be fed, including the +// stream adapters, and reports each result that differs or fails via t.Error or t.Errorf. +func Conformance(t testing.TB, newDecoder stream.NewDecoderFunc, c Case) { + t.Helper() + unwrap := c.Unwrap + if unwrap == nil { + unwrap = asDataSet + } + outcomes := make([]outcome, 0, len(feeds)+2) + for _, f := range feeds { + ts, err := decodeWith(newDecoder, c.TRQ, f.run, c.Body) + outcomes = append(outcomes, outcome{f.name, ts, err}) + } + ts, err := stream.ReaderUnmarshaler(newDecoder)(bytes.NewReader(c.Body), c.TRQ) + outcomes = append(outcomes, outcome{"reader-unmarshaler", ts, err}) + ts, err = stream.BytesUnmarshaler(newDecoder)(c.Body, c.TRQ) + outcomes = append(outcomes, outcome{"bytes-unmarshaler", ts, err}) + + if c.WantErr != nil { + for _, o := range outcomes { + if !errorMatches(o.err, c.WantErr) { + t.Errorf("%s: got error %v, want %v", o.name, o.err, c.WantErr) + } + } + if c.Legacy != nil { + if _, err := c.Legacy(bytes.NewReader(c.Body), c.TRQ); err == nil { + t.Error("legacy: got no error, want one") + } + } + return + } + + var base *dataset.DataSet + var baseName string + for _, o := range outcomes { + ds := checkOutcome(t, unwrap, o) + if ds == nil { + continue + } + if base == nil { + base, baseName = ds, o.name + continue + } + if err := Compare(base, ds, CompareOptions{}); err != nil { + t.Errorf("%s: differs from %s: %v", o.name, baseName, err) + } + } + if base == nil { + return + } + if c.Want != nil { + if err := Compare(c.Want, base, CompareOptions{IgnoreSizes: true}); err != nil { + t.Errorf("%s: differs from Want: %v", baseName, err) + } + } + if c.Legacy != nil { + ts, err := c.Legacy(bytes.NewReader(c.Body), c.TRQ) + if ds := checkOutcome(t, unwrap, outcome{"legacy", ts, err}); ds != nil { + if err := Compare(ds, base, CompareOptions{IgnoreSizes: true}); err != nil { + t.Errorf("%s: differs from legacy: %v", baseName, err) + } + } + } + if c.Shuffle != nil { + for i := range uint64(shuffleRounds) { + rng := weaktest.NewRand(i, 0x5eed) + ts, err := stream.BytesUnmarshaler(newDecoder)(c.Shuffle(slices.Clone(c.Body), rng), c.TRQ) + o := outcome{"shuffle-" + strconv.FormatUint(i, 10), ts, err} + if ds := checkOutcome(t, unwrap, o); ds != nil { + if err := Compare(base, ds, CompareOptions{IgnoreSeriesOrder: true}); err != nil { + t.Errorf("%s: differs from %s: %v", o.name, baseName, err) + } + } + } + } + // a failed read must surface as an error rather than as a partial result + r := io.MultiReader(bytes.NewReader(c.Body[:len(c.Body)/2]), iotest.ErrReader(errInjected)) + if _, err := stream.ReaderUnmarshaler(newDecoder)(r, c.TRQ); err == nil { + t.Error("read-error: got no error for a failed read") + } +} + +// ShuffleLines returns a Case.Shuffle that keeps the first header lines in place +// and shuffles the rest, for formats whose rows may arrive in any order. +func ShuffleLines(header int) func(body []byte, rng *weaktest.Rand) []byte { + return func(body []byte, rng *weaktest.Rand) []byte { + lines := bytes.SplitAfter(body, []byte{'\n'}) + if n := len(lines); n > 0 && len(lines[n-1]) == 0 { + lines = lines[:n-1] + } + if n := len(lines); n > 0 && !bytes.HasSuffix(lines[n-1], []byte{'\n'}) { + lines[n-1] = append(slices.Clip(lines[n-1]), '\n') + } + if header < len(lines) { + rows := lines[max(header, 0):] + rng.Shuffle(len(rows), func(i, j int) { rows[i], rows[j] = rows[j], rows[i] }) + } + return bytes.Join(lines, nil) + } +} + +// Bench measures u decoding body, as the proxy engine calls it, so a stream +// decoder and an existing unmarshaler can be compared. +func Bench(b *testing.B, u timeseries.UnmarshalerReaderFunc, + trq *timeseries.TimeRangeQuery, body []byte, +) { + b.Helper() + b.ReportAllocs() + b.SetBytes(int64(len(body))) + r := bytes.NewReader(body) + for b.Loop() { + r.Reset(body) + if _, err := u(io.NopCloser(r), trq); err != nil { + b.Fatal(err) + } + } +} + +func decodeWith(newDecoder stream.NewDecoderFunc, trq *timeseries.TimeRangeQuery, + run func(stream.Decoder, []byte) error, body []byte, +) (timeseries.Timeseries, error) { + dec, err := newDecoder(trq) + if err != nil { + return nil, err + } + if err := run(dec, body); err != nil { + return nil, err + } + return dec.Finish() +} + +func checkOutcome(t testing.TB, unwrap func(timeseries.Timeseries) *dataset.DataSet, + o outcome, +) *dataset.DataSet { + t.Helper() + if o.err != nil { + t.Errorf("%s: unexpected error: %v", o.name, o.err) + return nil + } + ds := unwrap(o.ts) + if ds == nil { + t.Errorf("%s: got %T, want a DataSet", o.name, o.ts) + } + return ds +} + +func asDataSet(ts timeseries.Timeseries) *dataset.DataSet { + ds, _ := ts.(*dataset.DataSet) + return ds +} + +func errorMatches(err, want error) bool { + if err == nil { + return false + } + return errors.Is(want, ErrAny) || errors.Is(err, want) +} diff --git a/pkg/timeseries/dataset/stream/streamtest/streamtest_test.go b/pkg/timeseries/dataset/stream/streamtest/streamtest_test.go new file mode 100644 index 000000000..fd60b64e5 --- /dev/null +++ b/pkg/timeseries/dataset/stream/streamtest/streamtest_test.go @@ -0,0 +1,293 @@ +/* + * Copyright 2018 The Trickster Authors + * + * 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 streamtest + +import ( + "errors" + "flag" + "fmt" + "io" + "strconv" + "strings" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/timeseries" + "github.com/trickstercache/trickster/v2/pkg/timeseries/dataset" + "github.com/trickstercache/trickster/v2/pkg/timeseries/dataset/stream" + "github.com/trickstercache/trickster/v2/pkg/timeseries/epoch" + "github.com/trickstercache/trickster/v2/pkg/util/weak/weaktest" + + "github.com/stretchr/testify/require" +) + +type recorder struct { + testing.TB + errs []string +} + +func (r *recorder) Helper() {} + +func (r *recorder) Errorf(format string, args ...any) { + r.errs = append(r.errs, fmt.Sprintf(format, args...)) +} + +func (r *recorder) Error(args ...any) { + r.errs = append(r.errs, fmt.Sprint(args...)) +} + +func (r *recorder) requireReported(t *testing.T, subs ...string) { + t.Helper() + require.NotEmpty(t, r.errs) + all := strings.Join(r.errs, "\n") + for _, sub := range subs { + require.Contains(t, all, sub) + } +} + +var ( + testTRQ = ×eries.TimeRangeQuery{ + Extent: timeseries.Extent{Start: time.Unix(0, 0), End: time.Unix(10, 0)}, + } + csvFields = timeseries.SeriesFields{ + Tags: timeseries.FieldDefinitions{{Name: "host", DataType: timeseries.String}}, + Values: timeseries.FieldDefinitions{{Name: "v", DataType: timeseries.String}}, + } + errPicky = errors.New("single byte write") + csvBody = []byte("3,a,x\n1,b,y\n2,a,z\n1,a,w\n") +) + +func csvDecoder(opts dataset.BuilderOptions) stream.NewDecoderFunc { + opts.Fields = csvFields + return func(trq *timeseries.TimeRangeQuery) (stream.Decoder, error) { + b := dataset.NewBuilder(trq, opts) + var cols [][]byte + onLine := func(line []byte) error { + cols = stream.SplitFields(line, ',', cols) + if len(cols) != 3 { + return timeseries.ErrInvalidBody + } + ep, err := epoch.ParseDecimal(cols[0], timeseries.DateTimeUnixSecs) + if err != nil { + return err + } + r := b.Row() + r.SetEpoch(ep) + r.SetTag(0, cols[1]) + r.AddValue(string(cols[2])) + return r.Commit() + } + return stream.NewLines(onLine, func() (timeseries.Timeseries, error) { + return b.Finish() + }), nil + } +} + +func wrapDecoder(newDecoder stream.NewDecoderFunc, + wrap func(stream.Decoder) stream.Decoder, +) stream.NewDecoderFunc { + return func(trq *timeseries.TimeRangeQuery) (stream.Decoder, error) { + dec, err := newDecoder(trq) + if err != nil { + return nil, err + } + return wrap(dec), nil + } +} + +type countingDecoder struct { + stream.Decoder + writes int +} + +func (c *countingDecoder) Write(p []byte) (int, error) { + c.writes++ + return c.Decoder.Write(p) +} + +func (c *countingDecoder) Finish() (timeseries.Timeseries, error) { + ts, err := c.Decoder.Finish() + if ds, ok := ts.(*dataset.DataSet); ok { + ds.Status = strconv.Itoa(c.writes) + } + return ts, err +} + +type pickyDecoder struct{ stream.Decoder } + +func (p pickyDecoder) Write(b []byte) (int, error) { + if len(b) == 1 { + return 0, errPicky + } + return p.Decoder.Write(b) +} + +type lenientDecoder struct{ stream.Decoder } + +func (l lenientDecoder) ReadFrom(r io.Reader) (int64, error) { + b, _ := io.ReadAll(r) + n, err := l.Write(b) + return int64(n), err +} + +type wrapped struct{ *dataset.DataSet } + +type wrappingDecoder struct{ stream.Decoder } + +func (w wrappingDecoder) Finish() (timeseries.Timeseries, error) { + ts, err := w.Decoder.Finish() + if err != nil { + return nil, err + } + return wrapped{ts.(*dataset.DataSet)}, nil +} + +func legacyFrom(newDecoder stream.NewDecoderFunc) timeseries.UnmarshalerReaderFunc { + return stream.ReaderUnmarshaler(newDecoder) +} + +func TestConformancePasses(t *testing.T) { + dec := csvDecoder(dataset.BuilderOptions{SeriesName: "csv"}) + want, err := stream.BytesUnmarshaler(dec)(csvBody, testTRQ) + require.NoError(t, err) + c := Case{TRQ: testTRQ, Body: csvBody, Want: want.(*dataset.DataSet), + Legacy: legacyFrom(dec), Shuffle: ShuffleLines(0)} + Conformance(t, dec, c) + rec := &recorder{} + Conformance(rec, dec, c) + require.Empty(t, rec.errs) + + wrap := wrapDecoder(dec, func(d stream.Decoder) stream.Decoder { return wrappingDecoder{d} }) + Conformance(t, wrap, Case{TRQ: testTRQ, Body: csvBody, Legacy: legacyFrom(wrap), + Unwrap: func(ts timeseries.Timeseries) *dataset.DataSet { return ts.(wrapped).DataSet }}) +} + +func TestConformanceReportsDifferences(t *testing.T) { + dec := wrapDecoder(csvDecoder(dataset.BuilderOptions{}), + func(d stream.Decoder) stream.Decoder { return &countingDecoder{Decoder: d} }) + rec := &recorder{} + Conformance(rec, dec, Case{TRQ: testTRQ, Body: csvBody}) + rec.requireReported(t, "write-bytewise: differs from write-once: status") +} + +func TestConformanceReportsFeedErrors(t *testing.T) { + dec := wrapDecoder(csvDecoder(dataset.BuilderOptions{}), + func(d stream.Decoder) stream.Decoder { return pickyDecoder{d} }) + rec := &recorder{} + Conformance(rec, dec, Case{TRQ: testTRQ, Body: csvBody}) + rec.requireReported(t, "write-bytewise: unexpected error: single byte write") + + failing := func(*timeseries.TimeRangeQuery) (stream.Decoder, error) { return nil, errPicky } + rec = &recorder{} + Conformance(rec, failing, Case{TRQ: testTRQ, Body: csvBody}) + require.Len(t, rec.errs, len(feeds)+2) + Conformance(t, failing, Case{TRQ: testTRQ, Body: csvBody, WantErr: errPicky}) +} + +func TestConformanceWantErr(t *testing.T) { + good := csvDecoder(dataset.BuilderOptions{}) + rec := &recorder{} + Conformance(rec, good, Case{TRQ: testTRQ, Body: csvBody, WantErr: ErrAny, Legacy: legacyFrom(good)}) + rec.requireReported(t, "write-once: got error , want any error", "legacy: got no error") + + bad := []byte("1,a\n") + Conformance(t, good, Case{TRQ: testTRQ, Body: bad, WantErr: ErrAny, Legacy: legacyFrom(good)}) + Conformance(t, good, Case{TRQ: testTRQ, Body: bad, WantErr: timeseries.ErrInvalidBody}) + rec = &recorder{} + Conformance(rec, good, Case{TRQ: testTRQ, Body: bad, WantErr: timeseries.ErrInvalidTimeFormat}) + rec.requireReported(t, "want invalid time format") + + rec = &recorder{} + Conformance(rec, good, Case{TRQ: testTRQ, Body: bad}) + rec.requireReported(t, "readfrom: unexpected error") +} + +func TestConformanceWantAndLegacy(t *testing.T) { + dec := csvDecoder(dataset.BuilderOptions{}) + other := csvDecoder(dataset.BuilderOptions{SeriesName: "other"}) + want, err := stream.BytesUnmarshaler(other)(csvBody, testTRQ) + require.NoError(t, err) + rec := &recorder{} + Conformance(rec, dec, Case{TRQ: testTRQ, Body: csvBody, Want: want.(*dataset.DataSet), + Legacy: legacyFrom(other)}) + rec.requireReported(t, "differs from Want: results[0].series[0].name", + "differs from legacy: results[0].series[0].name") + + failingLegacy := func(io.Reader, *timeseries.TimeRangeQuery) (timeseries.Timeseries, error) { + return nil, errPicky + } + rec = &recorder{} + Conformance(rec, dec, Case{TRQ: testTRQ, Body: csvBody, Legacy: failingLegacy}) + rec.requireReported(t, "legacy: unexpected error: single byte write") +} + +func TestConformanceShuffle(t *testing.T) { + // with first-wins duplicates, the surviving value depends on row order + dec := csvDecoder(dataset.BuilderOptions{Duplicates: dataset.DuplicatesFirstWins}) + body := []byte("1,a,x\n1,a,y\n1,a,z\n1,a,w\n2,a,v\n") + rec := &recorder{} + Conformance(rec, dec, Case{TRQ: testTRQ, Body: body, Shuffle: ShuffleLines(0)}) + rec.requireReported(t, "differs from write-once: results[0].series[0].points[0].values[0]") + + failing := func(body []byte, _ *weaktest.Rand) []byte { return append(body, "x\n"...) } + rec = &recorder{} + Conformance(rec, dec, Case{TRQ: testTRQ, Body: csvBody, Shuffle: failing}) + rec.requireReported(t, "shuffle-0: unexpected error") +} + +func TestConformanceReadError(t *testing.T) { + dec := wrapDecoder(csvDecoder(dataset.BuilderOptions{}), + func(d stream.Decoder) stream.Decoder { return lenientDecoder{d} }) + rec := &recorder{} + Conformance(rec, dec, Case{TRQ: testTRQ, Body: csvBody}) + rec.requireReported(t, "read-error: got no error for a failed read") +} + +func TestConformanceUnwrap(t *testing.T) { + dec := wrapDecoder(csvDecoder(dataset.BuilderOptions{}), + func(d stream.Decoder) stream.Decoder { return wrappingDecoder{d} }) + rec := &recorder{} + Conformance(rec, dec, Case{TRQ: testTRQ, Body: csvBody}) + rec.requireReported(t, "write-once: got streamtest.wrapped, want a DataSet") +} + +func TestShuffleLines(t *testing.T) { + rng := weaktest.NewRand(1, 2) + body := []byte("h1\nh2\na\nb\nc\nd\ne\nf") + for range 8 { + out := string(ShuffleLines(2)(body, rng)) + require.True(t, strings.HasPrefix(out, "h1\nh2\n")) + require.True(t, strings.HasSuffix(out, "\n")) + require.ElementsMatch(t, strings.Split("h1\nh2\na\nb\nc\nd\ne\nf", "\n"), + strings.Split(strings.TrimSuffix(out, "\n"), "\n")) + } + require.Equal(t, "a\nb\n", string(ShuffleLines(5)([]byte("a\nb\n"), rng))) + require.Empty(t, ShuffleLines(-1)(nil, rng)) + require.Equal(t, "a\n", string(ShuffleLines(-1)([]byte("a"), rng))) +} + +func TestBench(t *testing.T) { + bt := flag.Lookup("test.benchtime") + require.NotNil(t, bt) + prev := bt.Value.String() + require.NoError(t, bt.Value.Set("3x")) + t.Cleanup(func() { bt.Value.Set(prev) }) + u := stream.ReaderUnmarshaler(csvDecoder(dataset.BuilderOptions{})) + res := testing.Benchmark(func(b *testing.B) { Bench(b, u, testTRQ, csvBody) }) + require.Positive(t, res.N) + require.Equal(t, int64(len(csvBody)), res.Bytes) +} diff --git a/pkg/timeseries/dataset/stream/value.go b/pkg/timeseries/dataset/stream/value.go new file mode 100644 index 000000000..e59827a0e --- /dev/null +++ b/pkg/timeseries/dataset/stream/value.go @@ -0,0 +1,174 @@ +/* + * Copyright 2018 The Trickster Authors + * + * 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 stream + +import ( + "bytes" + "encoding/json" + "fmt" + "strconv" + + "github.com/trickstercache/trickster/v2/pkg/timeseries" +) + +// ErrInvalidValue indicates a value that cannot be parsed as its field's data type. +var ErrInvalidValue = fmt.Errorf("%w: invalid value", timeseries.ErrInvalidBody) + +var nullLiteral = []byte("null") + +// ParseValue parses the text of a value, such as a TSV or CSV cell, as dt. Text +// types return a string; other types return nil for empty text. +func ParseValue(raw []byte, dt timeseries.FieldDataType) (any, error) { + switch dt { + case timeseries.String, timeseries.DateTimeRFC3339, timeseries.DateTimeRFC3339Nano, + timeseries.DateSQL, timeseries.TimeSQL, timeseries.DateTimeSQL: + return string(raw), nil + case timeseries.Null: + return nil, nil + } + if len(raw) == 0 { + return nil, nil + } + var v any + var err error + switch dt { + case timeseries.Int64, timeseries.DateTimeUnixSecs, timeseries.DateTimeUnixMilli, + timeseries.DateTimeUnixMicro, timeseries.DateTimeUnixNano: + v, err = strconv.ParseInt(string(raw), 10, 64) + case timeseries.Int16: + v, err = strconv.ParseInt(string(raw), 10, 16) + case timeseries.Byte: + v, err = strconv.ParseInt(string(raw), 10, 8) + case timeseries.Uint64: + v, err = strconv.ParseUint(string(raw), 10, 64) + case timeseries.Float64: + v, err = strconv.ParseFloat(string(raw), 64) + case timeseries.Bool: + v, err = strconv.ParseBool(string(raw)) + case timeseries.Unknown: + return inferValue(raw), nil + default: + return nil, ErrInvalidValue + } + if err != nil { + return nil, fmt.Errorf("%w: %w", ErrInvalidValue, err) + } + return v, nil +} + +// ParseJSONValue parses a raw JSON value as dt: null is nil, and quoted values +// are unquoted before parsing, since some APIs send numbers as JSON strings. +func ParseJSONValue(raw []byte, dt timeseries.FieldDataType) (any, error) { + if bytes.Equal(raw, nullLiteral) { + return nil, nil + } + if !isQuoted(raw) { + return ParseValue(raw, dt) + } + b, err := unquote(raw) + if err != nil { + return nil, err + } + if dt == timeseries.Unknown { + return string(b), nil + } + return ParseValue(b, dt) +} + +func isQuoted(raw []byte) bool { + n := len(raw) + return n >= 2 && raw[0] == '"' && raw[n-1] == '"' +} + +func unquote(raw []byte) ([]byte, error) { + inner := raw[1 : len(raw)-1] + if bytes.IndexByte(inner, '\\') < 0 { + return inner, nil + } + var s string + if err := json.Unmarshal(raw, &s); err != nil { + return nil, fmt.Errorf("%w: %w", ErrInvalidValue, err) + } + return []byte(s), nil +} + +func inferValue(raw []byte) any { + // JSON literals become bools and numbers; any other text stays a string + switch string(raw) { + case "true": + return true + case "false": + return false + } + isInt, ok := jsonNumber(raw) + if ok && isInt { + if i, err := strconv.ParseInt(string(raw), 10, 64); err == nil { + return i + } + if u, err := strconv.ParseUint(string(raw), 10, 64); err == nil { + return u + } + } + if ok { + if f, err := strconv.ParseFloat(string(raw), 64); err == nil { + return f + } + } + return string(raw) +} + +func jsonNumber(b []byte) (isInt, ok bool) { + // ok reports a match of the JSON number grammar; isInt, the absence of a fraction and exponent + i := 0 + if i < len(b) && b[i] == '-' { + i++ + } + start := i + for i < len(b) && b[i] >= '0' && b[i] <= '9' { + i++ + } + if i == start || (b[start] == '0' && i-start > 1) { + return false, false + } + isInt = true + if i < len(b) && b[i] == '.' { + isInt = false + i++ + start = i + for i < len(b) && b[i] >= '0' && b[i] <= '9' { + i++ + } + if i == start { + return false, false + } + } + if i < len(b) && (b[i] == 'e' || b[i] == 'E') { + isInt = false + i++ + if i < len(b) && (b[i] == '+' || b[i] == '-') { + i++ + } + start = i + for i < len(b) && b[i] >= '0' && b[i] <= '9' { + i++ + } + if i == start { + return false, false + } + } + return isInt, i == len(b) +} diff --git a/pkg/timeseries/dataset/stream/value_test.go b/pkg/timeseries/dataset/stream/value_test.go new file mode 100644 index 000000000..c61d7befc --- /dev/null +++ b/pkg/timeseries/dataset/stream/value_test.go @@ -0,0 +1,149 @@ +/* + * Copyright 2018 The Trickster Authors + * + * 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 stream + +import ( + "math" + "testing" + + "github.com/trickstercache/trickster/v2/pkg/timeseries" + + "github.com/stretchr/testify/require" +) + +func TestParseValue(t *testing.T) { + tests := []struct { + raw string + dt timeseries.FieldDataType + want any + }{ + {"abc", timeseries.String, "abc"}, + {"", timeseries.String, ""}, + {"null", timeseries.String, "null"}, + {"2024-01-01", timeseries.DateSQL, "2024-01-01"}, + {"x", timeseries.Null, nil}, + {"", timeseries.Int64, nil}, + {"-42", timeseries.Int64, int64(-42)}, + {"1727222400", timeseries.DateTimeUnixSecs, int64(1727222400)}, + {"-300", timeseries.Int16, int64(-300)}, + {"-100", timeseries.Byte, int64(-100)}, + {"18446744073709551615", timeseries.Uint64, uint64(math.MaxUint64)}, + {"1.5", timeseries.Float64, 1.5}, + {"NaN", timeseries.Float64, math.NaN()}, + {"-Inf", timeseries.Float64, math.Inf(-1)}, + {"true", timeseries.Bool, true}, + {"0", timeseries.Bool, false}, + {"true", timeseries.Unknown, true}, + {"false", timeseries.Unknown, false}, + {"-12", timeseries.Unknown, int64(-12)}, + {"18446744073709551615", timeseries.Unknown, uint64(math.MaxUint64)}, + {"99999999999999999999", timeseries.Unknown, 1e20}, + {"-0.5", timeseries.Unknown, -0.5}, + {"1e3", timeseries.Unknown, 1000.0}, + {"1E+3", timeseries.Unknown, 1000.0}, + {"2.5e-1", timeseries.Unknown, 0.25}, + {"0", timeseries.Unknown, int64(0)}, + {"1e999", timeseries.Unknown, "1e999"}, + {"01", timeseries.Unknown, "01"}, + {"-", timeseries.Unknown, "-"}, + {"1.", timeseries.Unknown, "1."}, + {"1e", timeseries.Unknown, "1e"}, + {"1e+", timeseries.Unknown, "1e+"}, + {"1x", timeseries.Unknown, "1x"}, + {"NaN", timeseries.Unknown, "NaN"}, + {"abc", timeseries.Unknown, "abc"}, + } + for _, test := range tests { + t.Run(test.raw, func(t *testing.T) { + got, err := ParseValue([]byte(test.raw), test.dt) + require.NoError(t, err) + if f, ok := test.want.(float64); ok && math.IsNaN(f) { + require.True(t, math.IsNaN(got.(float64))) + return + } + require.Equal(t, test.want, got) + }) + } +} + +func TestParseValueErrors(t *testing.T) { + tests := []struct { + raw string + dt timeseries.FieldDataType + }{ + {"x", timeseries.Int64}, + {"1.5", timeseries.Int64}, + {"40000", timeseries.Int16}, + {"200", timeseries.Byte}, + {"-1", timeseries.Uint64}, + {"x", timeseries.Float64}, + {"yes", timeseries.Bool}, + {"1", timeseries.FieldDataType(250)}, + } + for _, test := range tests { + t.Run(test.raw, func(t *testing.T) { + _, err := ParseValue([]byte(test.raw), test.dt) + require.ErrorIs(t, err, ErrInvalidValue) + require.ErrorIs(t, err, timeseries.ErrInvalidBody) + }) + } +} + +func TestParseJSONValue(t *testing.T) { + tests := []struct { + raw string + dt timeseries.FieldDataType + want any + }{ + {"null", timeseries.String, nil}, + {"null", timeseries.Float64, nil}, + {`"null"`, timeseries.String, "null"}, + {`"abc"`, timeseries.String, "abc"}, + {`"a\"b\\c\n"`, timeseries.String, "a\"b\\c\n"}, + {`""`, timeseries.String, ""}, + {`""`, timeseries.Float64, nil}, + {`"1.5"`, timeseries.Float64, 1.5}, + {`1.5`, timeseries.Float64, 1.5}, + {`"+Inf"`, timeseries.Float64, math.Inf(1)}, + {`123`, timeseries.String, "123"}, + {`"123"`, timeseries.Unknown, "123"}, + {`123`, timeseries.Unknown, int64(123)}, + {`true`, timeseries.Unknown, true}, + {`"`, timeseries.Unknown, `"`}, + } + for _, test := range tests { + t.Run(test.raw, func(t *testing.T) { + got, err := ParseJSONValue([]byte(test.raw), test.dt) + require.NoError(t, err) + require.Equal(t, test.want, got) + }) + } + _, err := ParseJSONValue([]byte(`"\x"`), timeseries.String) + require.ErrorIs(t, err, ErrInvalidValue) + _, err = ParseJSONValue([]byte(`"x"`), timeseries.Int64) + require.ErrorIs(t, err, ErrInvalidValue) +} + +func BenchmarkParseJSONValue(b *testing.B) { + raw := []byte("1727222400.123") + b.ReportAllocs() + for b.Loop() { + if _, err := ParseJSONValue(raw, timeseries.Float64); err != nil { + b.Fatal(err) + } + } +} diff --git a/pkg/timeseries/epoch/parse.go b/pkg/timeseries/epoch/parse.go new file mode 100644 index 000000000..0bb6939a9 --- /dev/null +++ b/pkg/timeseries/epoch/parse.go @@ -0,0 +1,136 @@ +/* + * Copyright 2018 The Trickster Authors + * + * 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 epoch + +import ( + "math" + + "github.com/trickstercache/trickster/v2/pkg/timeseries" +) + +// ParseDecimal parses a base-10 count of the unit (DateTimeUnixSecs, Milli, Micro +// or Nano) into an Epoch with integer math, rejecting precision finer than 1ns. +func ParseDecimal(raw []byte, unit timeseries.FieldDataType) (Epoch, error) { + var exp int + switch unit { + case timeseries.DateTimeUnixSecs: + exp = 9 + case timeseries.DateTimeUnixMilli: + exp = 6 + case timeseries.DateTimeUnixMicro: + exp = 3 + case timeseries.DateTimeUnixNano: + default: + return 0, timeseries.ErrInvalidTimeFormat + } + b := raw + if n := len(b); n >= 2 && b[0] == '"' && b[n-1] == '"' { + b = b[1 : n-1] + } + var neg bool + if len(b) > 0 && (b[0] == '-' || b[0] == '+') { + neg = b[0] == '-' + b = b[1:] + } + intPart, b := leadingDigits(b) + var frac []byte + if len(b) > 0 && b[0] == '.' { + frac, b = leadingDigits(b[1:]) + } + if len(intPart)+len(frac) == 0 { + return 0, timeseries.ErrInvalidTimeFormat + } + if len(b) > 0 && (b[0] == 'e' || b[0] == 'E') { + e, ok := parseExponent(b[1:]) + if !ok { + return 0, timeseries.ErrInvalidTimeFormat + } + exp += e + } else if len(b) > 0 { + return 0, timeseries.ErrInvalidTimeFormat + } + // trailing fractional zeros carry no value and would only risk overflow + for len(frac) > 0 && frac[len(frac)-1] == '0' { + frac = frac[:len(frac)-1] + } + exp -= len(frac) + var v int64 + for _, digits := range [2][]byte{intPart, frac} { + for _, c := range digits { + d := int64(c - '0') + if v > (math.MaxInt64-d)/10 { + return 0, timeseries.ErrInvalidTimeFormat + } + v = v*10 + d + } + } + if v == 0 { + return 0, nil + } + // a nonzero int64 has at most 19 digits, so larger shifts cannot be exact + if exp > 19 || exp < -19 { + return 0, timeseries.ErrInvalidTimeFormat + } + for ; exp > 0; exp-- { + if v > math.MaxInt64/10 { + return 0, timeseries.ErrInvalidTimeFormat + } + v *= 10 + } + for ; exp < 0; exp++ { + if v%10 != 0 { + return 0, timeseries.ErrInvalidTimeFormat + } + v /= 10 + } + if neg { + v = -v + } + return Epoch(v), nil +} + +func leadingDigits(b []byte) (digits, rest []byte) { + i := 0 + for i < len(b) && b[i] >= '0' && b[i] <= '9' { + i++ + } + return b[:i], b[i:] +} + +func parseExponent(raw []byte) (int, bool) { + b := raw + var neg bool + if len(b) > 0 && (b[0] == '-' || b[0] == '+') { + neg = b[0] == '-' + b = b[1:] + } + digits, rest := leadingDigits(b) + if len(digits) == 0 || len(rest) > 0 { + return 0, false + } + var e int + for _, c := range digits { + // any exponent beyond this is rejected by the caller's range check + if e < 1000 { + e = e*10 + int(c-'0') + } + } + if neg { + e = -e + } + return e, true +} diff --git a/pkg/timeseries/epoch/parse_test.go b/pkg/timeseries/epoch/parse_test.go new file mode 100644 index 000000000..3b27f9ef1 --- /dev/null +++ b/pkg/timeseries/epoch/parse_test.go @@ -0,0 +1,109 @@ +/* + * Copyright 2018 The Trickster Authors + * + * 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 epoch + +import ( + "math" + "strconv" + "testing" + + "github.com/trickstercache/trickster/v2/pkg/timeseries" + + "github.com/stretchr/testify/require" +) + +func TestParseDecimal(t *testing.T) { + tests := []struct { + in string + unit timeseries.FieldDataType + want Epoch + }{ + {"1727222400", timeseries.DateTimeUnixSecs, 1727222400 * 1e9}, + {"1727222400.123", timeseries.DateTimeUnixSecs, 1727222400123000000}, + {"1727222400.123456789", timeseries.DateTimeUnixSecs, 1727222400123456789}, + {"1727222400.1230000000000000000000", timeseries.DateTimeUnixSecs, 1727222400123000000}, + {"1.727222400123e+09", timeseries.DateTimeUnixSecs, 1727222400123000000}, + {"1.727222400123E9", timeseries.DateTimeUnixSecs, 1727222400123000000}, + {"1727222400123e-3", timeseries.DateTimeUnixSecs, 1727222400123000000}, + {"1000e-12", timeseries.DateTimeUnixSecs, 1}, + {`"1727222400.5"`, timeseries.DateTimeUnixSecs, 1727222400500000000}, + {".5", timeseries.DateTimeUnixSecs, 500000000}, + {"5.", timeseries.DateTimeUnixSecs, 5000000000}, + {"+5", timeseries.DateTimeUnixSecs, 5000000000}, + {"-1.5", timeseries.DateTimeUnixSecs, -1500000000}, + {"0", timeseries.DateTimeUnixSecs, 0}, + {"0.000", timeseries.DateTimeUnixSecs, 0}, + {"0e99999", timeseries.DateTimeUnixSecs, 0}, + {"000000000000000000000001", timeseries.DateTimeUnixNano, 1}, + {"1727222400123", timeseries.DateTimeUnixMilli, 1727222400123000000}, + {"1727222400123.456", timeseries.DateTimeUnixMilli, 1727222400123456000}, + {"1727222400123456", timeseries.DateTimeUnixMicro, 1727222400123456000}, + {"1727222400123456789", timeseries.DateTimeUnixNano, 1727222400123456789}, + {strconv.FormatInt(math.MaxInt64, 10), timeseries.DateTimeUnixNano, math.MaxInt64}, + } + for _, test := range tests { + t.Run(test.in, func(t *testing.T) { + got, err := ParseDecimal([]byte(test.in), test.unit) + require.NoError(t, err) + require.Equal(t, test.want, got) + }) + } +} + +func TestParseDecimalErrors(t *testing.T) { + tests := []struct { + in string + unit timeseries.FieldDataType + }{ + {"1", timeseries.Float64}, + {"", timeseries.DateTimeUnixSecs}, + {`""`, timeseries.DateTimeUnixSecs}, + {"-", timeseries.DateTimeUnixSecs}, + {".", timeseries.DateTimeUnixSecs}, + {"e5", timeseries.DateTimeUnixSecs}, + {"1e", timeseries.DateTimeUnixSecs}, + {"1e+", timeseries.DateTimeUnixSecs}, + {"1e5x", timeseries.DateTimeUnixSecs}, + {"1x", timeseries.DateTimeUnixSecs}, + {"1.2.3", timeseries.DateTimeUnixSecs}, + {" 1", timeseries.DateTimeUnixSecs}, + {"0.0000000001", timeseries.DateTimeUnixSecs}, + {"1.5", timeseries.DateTimeUnixNano}, + {"1e-20", timeseries.DateTimeUnixNano}, + {"1e20", timeseries.DateTimeUnixNano}, + {"1e99999", timeseries.DateTimeUnixSecs}, + {"9223372036854775808", timeseries.DateTimeUnixNano}, + {"9223372036854775807", timeseries.DateTimeUnixSecs}, + {"99999999999999999999", timeseries.DateTimeUnixNano}, + } + for _, test := range tests { + t.Run(test.in, func(t *testing.T) { + _, err := ParseDecimal([]byte(test.in), test.unit) + require.ErrorIs(t, err, timeseries.ErrInvalidTimeFormat) + }) + } +} + +func BenchmarkParseDecimal(b *testing.B) { + in := []byte("1727222400.123") + b.ReportAllocs() + for b.Loop() { + if _, err := ParseDecimal(in, timeseries.DateTimeUnixSecs); err != nil { + b.Fatal(err) + } + } +} diff --git a/pkg/timeseries/request_options.go b/pkg/timeseries/request_options.go index 268313be7..fb5578ba3 100644 --- a/pkg/timeseries/request_options.go +++ b/pkg/timeseries/request_options.go @@ -45,6 +45,10 @@ type RequestOptions struct { // wire body, including when DPC marshals into a buffer rather than a response writer. ResponseContentType string ResponseContentEncoding string + // FallbackToProxyOnError retries the original read-only query when any + // extent fails or the merged response cannot be modeled faithfully. + // Providers enabling this must ensure that replaying the query is safe. + FallbackToProxyOnError bool } // ExtractFastForwardDisabled will look for the FastForwardUserDisableFlag in the provided string diff --git a/pkg/util/middleware/mirror.go b/pkg/util/middleware/mirror.go index 275811951..f91c44f91 100644 --- a/pkg/util/middleware/mirror.go +++ b/pkg/util/middleware/mirror.go @@ -20,7 +20,6 @@ import ( "bytes" "context" "io" - "math/rand/v2" "net/http" "sync/atomic" @@ -33,6 +32,7 @@ import ( "github.com/trickstercache/trickster/v2/pkg/proxy/methods" po "github.com/trickstercache/trickster/v2/pkg/proxy/paths/options" "github.com/trickstercache/trickster/v2/pkg/proxy/request" + "github.com/trickstercache/trickster/v2/pkg/util/weak/compat" "github.com/prometheus/client_golang/prometheus" ) @@ -67,7 +67,7 @@ func Mirror(backendName string, o *po.MirrorOptions, target backends.Backend, ne } return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { // sampling a share of traffic needs no unpredictability - if !tctx.IsMirrored(r.Context()) && (m.percent >= 100 || rand.IntN(100) < m.percent) { //nolint:gosec // traffic sampling + if !tctx.IsMirrored(r.Context()) && (m.percent >= 100 || compat.IntN(100) < m.percent) { m.fire(r) } next.ServeHTTP(w, r) diff --git a/pkg/util/middleware/simulated_latency.go b/pkg/util/middleware/simulated_latency.go index 23045798f..171fa2d38 100644 --- a/pkg/util/middleware/simulated_latency.go +++ b/pkg/util/middleware/simulated_latency.go @@ -17,10 +17,11 @@ package middleware import ( - "math/rand" "net/http" "strconv" "time" + + "github.com/trickstercache/trickster/v2/pkg/util/weak/compat" ) const latencyHeaderName = "x-simulated-latency" @@ -37,7 +38,7 @@ func processSimulatedLatency(w http.ResponseWriter, minLatency, maxLatency time. if minMS >= maxMS { ms = minMS } else { - ms = (rand.Int63()%maxMS - minMS) + minMS // #nosec G404 -- we are OK with a weak random source, random-ish enough for our purposes, no security risk + ms = (compat.Int64()%maxMS - minMS) + minMS } if ms <= 0 { return diff --git a/pkg/util/weak/compat/compat.go b/pkg/util/weak/compat/compat.go new file mode 100644 index 000000000..8c56ed359 --- /dev/null +++ b/pkg/util/weak/compat/compat.go @@ -0,0 +1,39 @@ +/* + * Copyright 2018 The Trickster Authors + * + * 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 compat collects the application's non-cryptographic randomness in one +// place, so each use can be reviewed and replaced when it is no longer suitable. +package compat + +import "math/rand/v2" + +// Uint64 returns a random uint64 from a randomly seeded source; it must not be +// used for secrets. +func Uint64() uint64 { + return rand.Uint64() // #nosec G404 -- callers use it for identifiers and spreading, not secrets +} + +// IntN returns a random int in [0, n) from a randomly seeded source; it must not +// be used for secrets. It panics if n <= 0. +func IntN(n int) int { + return rand.IntN(n) // #nosec G404 -- callers use it for sampling, not secrets +} + +// Int64 returns a random non-negative int64 from a randomly seeded source; it +// must not be used for secrets. +func Int64() int64 { + return rand.Int64() // #nosec G404 -- callers use it for jitter, not secrets +} diff --git a/pkg/util/weak/compat/compat_test.go b/pkg/util/weak/compat/compat_test.go new file mode 100644 index 000000000..c203f7aca --- /dev/null +++ b/pkg/util/weak/compat/compat_test.go @@ -0,0 +1,36 @@ +/* + * Copyright 2018 The Trickster Authors + * + * 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 compat + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestCompat(t *testing.T) { + seen := make(map[uint64]bool) + for range 100 { + seen[Uint64()] = true + n := IntN(3) + require.GreaterOrEqual(t, n, 0) + require.Less(t, n, 3) + require.GreaterOrEqual(t, Int64(), int64(0)) + } + require.Greater(t, len(seen), 1) + require.Panics(t, func() { IntN(0) }) +} diff --git a/pkg/util/weak/export_test.go b/pkg/util/weak/export_test.go new file mode 100644 index 000000000..b4555e6fc --- /dev/null +++ b/pkg/util/weak/export_test.go @@ -0,0 +1,22 @@ +/* + * Copyright 2018 The Trickster Authors + * + * 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 weak + +// ResetNonTestUsage undoes RegisterNonTestUsage; it exists only in test builds. +func ResetNonTestUsage() { + nonTestUsage.Store(false) +} diff --git a/pkg/util/weak/weak.go b/pkg/util/weak/weak.go new file mode 100644 index 000000000..016f7012a --- /dev/null +++ b/pkg/util/weak/weak.go @@ -0,0 +1,52 @@ +/* + * Copyright 2018 The Trickster Authors + * + * 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 weak separates the application's non-cryptographic randomness, in +// package compat, from the reproducible randomness tests use, in package weaktest. +package weak + +import ( + "os" + "sync/atomic" +) + +// TestModeEnv names the environment variable that tests set, to any nonempty +// value, when they launch a trickster process that should stay in test mode. +const TestModeEnv = "TRICKSTER_TEST_MODE" + +var nonTestUsage atomic.Bool + +// RegisterNonTestUsage marks the process as the application rather than a test, +// after which package weaktest panics when used. +func RegisterNonTestUsage() { + nonTestUsage.Store(true) +} + +// RegisterUnlessTestMode calls RegisterNonTestUsage unless TestModeEnv is set. +// The application's main calls it first. +func RegisterUnlessTestMode() { + if os.Getenv(TestModeEnv) == "" { + RegisterNonTestUsage() + } +} + +// AssertTestUsage panics if RegisterNonTestUsage has been called. Test-only +// randomness calls it so that the application cannot use it. +func AssertTestUsage() { + if nonTestUsage.Load() { + panic("weak: test-only randomness used outside of a test; use package compat") + } +} diff --git a/pkg/util/weak/weak_test.go b/pkg/util/weak/weak_test.go new file mode 100644 index 000000000..9061893c0 --- /dev/null +++ b/pkg/util/weak/weak_test.go @@ -0,0 +1,56 @@ +/* + * Copyright 2018 The Trickster Authors + * + * 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 weak_test + +import ( + "testing" + + "github.com/trickstercache/trickster/v2/pkg/util/weak" + "github.com/trickstercache/trickster/v2/pkg/util/weak/compat" + "github.com/trickstercache/trickster/v2/pkg/util/weak/weaktest" + + "github.com/stretchr/testify/require" +) + +func TestRegisterUnlessTestMode(t *testing.T) { + t.Cleanup(weak.ResetNonTestUsage) + require.NotPanics(t, weak.AssertTestUsage) + + t.Setenv(weak.TestModeEnv, "1") + weak.RegisterUnlessTestMode() + require.NotPanics(t, weak.AssertTestUsage) + + t.Setenv(weak.TestModeEnv, "") + weak.RegisterUnlessTestMode() + require.Panics(t, weak.AssertTestUsage) +} + +func TestTestOnlyRandomnessPanicsInApplication(t *testing.T) { + t.Cleanup(weak.ResetNonTestUsage) + require.NotPanics(t, func() { weaktest.NewRand(1, 2).Uint64() }) + require.NotPanics(t, func() { weaktest.IntN(2) }) + + weak.RegisterNonTestUsage() + require.PanicsWithValue(t, "weak: test-only randomness used outside of a test; use package compat", + func() { weaktest.NewRand(1, 2) }) + require.Panics(t, func() { weaktest.IntN(2) }) + require.NotPanics(t, func() { + compat.Uint64() + compat.IntN(2) + compat.Int64() + }) +} diff --git a/pkg/util/weak/weaktest/weaktest.go b/pkg/util/weak/weaktest/weaktest.go new file mode 100644 index 000000000..4634e3a05 --- /dev/null +++ b/pkg/util/weak/weaktest/weaktest.go @@ -0,0 +1,42 @@ +/* + * Copyright 2018 The Trickster Authors + * + * 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 weaktest provides non-cryptographic randomness for tests and test +// helpers. Every function panics once the application has registered at startup. +package weaktest + +import ( + "math/rand/v2" + + "github.com/trickstercache/trickster/v2/pkg/util/weak" +) + +// Rand is a pseudo-random number generator returned by NewRand. +type Rand = rand.Rand + +// NewRand returns a Rand whose sequence is fixed by seed1 and seed2, so test data +// and shuffles are reproducible. +func NewRand(seed1, seed2 uint64) *Rand { + weak.AssertTestUsage() + return rand.New(rand.NewPCG(seed1, seed2)) // #nosec G404 -- test-only; weak.AssertTestUsage keeps it out of the application +} + +// IntN returns a random int in [0, n) from a randomly seeded source. It panics +// if n <= 0. +func IntN(n int) int { + weak.AssertTestUsage() + return rand.IntN(n) // #nosec G404 -- test-only; weak.AssertTestUsage keeps it out of the application +} diff --git a/pkg/util/weak/weaktest/weaktest_test.go b/pkg/util/weak/weaktest/weaktest_test.go new file mode 100644 index 000000000..ce96d4497 --- /dev/null +++ b/pkg/util/weak/weaktest/weaktest_test.go @@ -0,0 +1,43 @@ +/* + * Copyright 2018 The Trickster Authors + * + * 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 weaktest + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestNewRandIsReproducible(t *testing.T) { + a, b, c := NewRand(7, 11), NewRand(7, 11), NewRand(7, 12) + var differs bool + for range 8 { + x := a.Uint64() + require.Equal(t, x, b.Uint64()) + differs = differs || x != c.Uint64() + } + require.True(t, differs) +} + +func TestIntN(t *testing.T) { + for range 100 { + n := IntN(3) + require.GreaterOrEqual(t, n, 0) + require.Less(t, n, 3) + } + require.Panics(t, func() { IntN(0) }) +}