diff --git a/Cargo.lock b/Cargo.lock index ba3f508486e..de3bc7ca765 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -8556,7 +8556,7 @@ dependencies = [ [[package]] name = "meter-core" version = "0.1.0" -source = "git+https://github.com/GreptimeTeam/greptime-meter.git?rev=5618e779cf2bb4755b499c630fba4c35e91898cb#5618e779cf2bb4755b499c630fba4c35e91898cb" +source = "git+https://github.com/GreptimeTeam/greptime-meter.git?rev=f743c3a5fe8c1f57363eac5866c356b7d91f9b39#f743c3a5fe8c1f57363eac5866c356b7d91f9b39" dependencies = [ "anymap2", "once_cell", @@ -8567,7 +8567,7 @@ dependencies = [ [[package]] name = "meter-macros" version = "0.1.0" -source = "git+https://github.com/GreptimeTeam/greptime-meter.git?rev=5618e779cf2bb4755b499c630fba4c35e91898cb#5618e779cf2bb4755b499c630fba4c35e91898cb" +source = "git+https://github.com/GreptimeTeam/greptime-meter.git?rev=f743c3a5fe8c1f57363eac5866c356b7d91f9b39#f743c3a5fe8c1f57363eac5866c356b7d91f9b39" dependencies = [ "meter-core", ] @@ -13617,6 +13617,8 @@ dependencies = [ "local-ip-address", "log-query", "loki-proto", + "meter-core", + "meter-macros", "metric-engine", "mime_guess", "moka", diff --git a/Cargo.toml b/Cargo.toml index 27e829a8cd0..fa01426bbbf 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -178,7 +178,7 @@ jsonb = { git = "https://github.com/GreptimeTeam/jsonb", rev = "dd1e2341e3df35e9 lazy_static = "1.4" local-ip-address = "0.6" loki-proto = { git = "https://github.com/GreptimeTeam/loki-proto.git", rev = "f69c8924c4babe516373e26a4118be82d976629c" } -meter-core = { git = "https://github.com/GreptimeTeam/greptime-meter.git", rev = "5618e779cf2bb4755b499c630fba4c35e91898cb" } +meter-core = { git = "https://github.com/GreptimeTeam/greptime-meter.git", rev = "f743c3a5fe8c1f57363eac5866c356b7d91f9b39" } mockall = "0.13" moka = "0.12" nalgebra = "0.33" @@ -357,7 +357,7 @@ table = { path = "src/table" } [workspace.dependencies.meter-macros] git = "https://github.com/GreptimeTeam/greptime-meter.git" -rev = "5618e779cf2bb4755b499c630fba4c35e91898cb" +rev = "f743c3a5fe8c1f57363eac5866c356b7d91f9b39" [patch.crates-io] substrait = { git = "https://github.com/GreptimeTeam/substrait-rs.git", rev = "91ec978b0649417ad3da8390e7a515baec723b1b" } diff --git a/src/common/base/src/protocol.rs b/src/common/base/src/protocol.rs index 66704dbc18b..f9ade9630c5 100644 --- a/src/common/base/src/protocol.rs +++ b/src/common/base/src/protocol.rs @@ -14,7 +14,7 @@ use std::fmt::{Display, Formatter}; -/// The protocol through which a query is received. +/// The protocol or internal subsystem through which a query is received. #[derive(Debug, PartialEq, Default, Clone, Copy, strum::FromRepr)] #[repr(u8)] pub enum Channel { @@ -35,6 +35,8 @@ pub enum Channel { Log = 12, Promql = 13, Splunk = 14, + /// Trusted internal requests and local subsystem execution. + Internal = 255, } impl From for Channel { @@ -70,6 +72,7 @@ impl AsRef for Channel { Self::Log => "log", Self::Promql => "promql", Self::Splunk => "splunk", + Self::Internal => "internal", } } } @@ -95,6 +98,7 @@ mod tests { (12, "log"), (13, "promql"), (14, "splunk"), + (255, "internal"), ]; for (value, name) in expected { @@ -102,5 +106,6 @@ mod tests { } assert_eq!("unknown", Channel::from(0).as_ref()); assert_eq!("unknown", Channel::from(15).as_ref()); + assert_eq!("unknown", Channel::from(256).as_ref()); } } diff --git a/src/flow/AGENTS.md b/src/flow/AGENTS.md index 0b024e96fbd..266621f508f 100644 --- a/src/flow/AGENTS.md +++ b/src/flow/AGENTS.md @@ -54,6 +54,8 @@ default, and sink-creation helpers. ## Public surface - gRPC `Flow` service in `src/flow/src/server.rs`. +- Local streaming sink writes and batching query execution use `Channel::Internal` + so meter collectors can account for sink writes without charging ingestion quota. - `FlowEngine` trait in `src/flow/src/engine.rs`. - Started from the `cmd` crate via `FlownodeBuilder` / `FlownodeInstance`. diff --git a/src/flow/src/batching_mode/frontend_client.rs b/src/flow/src/batching_mode/frontend_client.rs index efec4cdea63..aa2f2aa072a 100644 --- a/src/flow/src/batching_mode/frontend_client.rs +++ b/src/flow/src/batching_mode/frontend_client.rs @@ -34,7 +34,7 @@ use query::options::{FlowQueryExtensions, QueryOptions}; use rand::rng; use rand::seq::SliceRandom; use servers::query_handler::grpc::GrpcQueryHandler; -use session::context::{QueryContextBuilder, QueryContextRef}; +use session::context::{Channel, QueryContextBuilder, QueryContextRef}; use session::hints::READ_PREFERENCE_HINT; use snafu::{OptionExt, ResultExt}; use tokio::sync::SetOnce; @@ -521,6 +521,7 @@ impl FrontendClient { .current_catalog(catalog.to_string()) .current_schema(schema.to_string()) .extensions(extensions_map) + .channel(Channel::Internal) .snapshot_seqs(Arc::new(RwLock::new(snapshot_seqs.clone()))) .build(); let ctx = Arc::new(ctx); @@ -589,6 +590,7 @@ impl FrontendClient { let ctx = QueryContextBuilder::default() .current_catalog(catalog.to_string()) .current_schema(schema.to_string()) + .channel(Channel::Internal) .extensions(HashMap::from([( QUERY_PARALLELISM_HINT.to_string(), query.parallelism.to_string(), @@ -894,8 +896,9 @@ mod tests { async fn do_query( &self, _query: Request, - _ctx: QueryContextRef, + ctx: QueryContextRef, ) -> std::result::Result { + assert_eq!(ctx.channel(), Channel::Internal); self.calls.fetch_add(1, Ordering::SeqCst); Ok(Output::new_with_affected_rows(1)) } @@ -941,6 +944,7 @@ mod tests { ctx: QueryContextRef, ) -> std::result::Result { assert_eq!(ctx.extension("flow.return_region_seq"), Some("true")); + assert_eq!(ctx.channel(), Channel::Internal); Ok(Output::new_with_affected_rows(1)) } } diff --git a/src/frontend/AGENTS.md b/src/frontend/AGENTS.md index 9740d160164..d552d73cf39 100644 --- a/src/frontend/AGENTS.md +++ b/src/frontend/AGENTS.md @@ -41,8 +41,15 @@ remote datanodes via `operator`/`client`. `StatementExecutor`. Distributed scans enter through `region_query.rs`. - **Insert** (`instance/grpc.rs`): `handle_inserts` / `handle_row_inserts` → `check_permission` → `operator`'s `Inserter` (schema validation, optional - auto-create, partition routing) → local `RegionServer` (standalone) or RPC to - datanodes (distributed). + auto-create, partition routing, meter admission) → local `RegionServer` + (standalone) or RPC to datanodes (distributed). Arrow bulk inserts pass the + request channel to `Inserter` and check meter admission for each nonempty batch. +- Finite ingestion requests split internally admit their total rows per database + before dispatch (`operator::insert::admit_write` / `admit_row_insert_batches`). + The returned context covers chunks and derived writes while preserving WCU + accounting and the original protocol channel. +- Internal gRPC listeners mark requests with `Channel::Internal` in middleware + (`server.rs`), including requests handled by Enterprise Flight wrappers. - **Table batching** (`instance/builder.rs`): protocol entry points opt in through `QueryContext`. The primary inserter prepares eligible ordinary-table writes diff --git a/src/frontend/src/instance/grpc.rs b/src/frontend/src/instance/grpc.rs index 2d89b5489df..e2435b65689 100644 --- a/src/frontend/src/instance/grpc.rs +++ b/src/frontend/src/instance/grpc.rs @@ -454,6 +454,7 @@ impl Instance { request.record_batch, request.schema_bytes, ctx.skip_wal(), + ctx.channel(), ) .await .context(TableOperationSnafu)?; diff --git a/src/frontend/src/instance/log_handler.rs b/src/frontend/src/instance/log_handler.rs index 051018c9a55..44738b9dcc1 100644 --- a/src/frontend/src/instance/log_handler.rs +++ b/src/frontend/src/instance/log_handler.rs @@ -80,6 +80,11 @@ impl PipelineHandler for Instance { prepared.push((Arc::new(ctx.fork()), log)); } + operator::insert::admit_row_insert_batches(&mut prepared) + .await + .map_err(BoxedError::new) + .context(servers::error::ExecuteGrpcQuerySnafu)?; + let mut outputs = Vec::with_capacity(prepared.len()); for (ctx, log) in prepared { outputs.push(self.handle_log_inserts(log, ctx).await); diff --git a/src/frontend/src/instance/otlp.rs b/src/frontend/src/instance/otlp.rs index d72b79a4275..fa40c7eb59a 100644 --- a/src/frontend/src/instance/otlp.rs +++ b/src/frontend/src/instance/otlp.rs @@ -30,6 +30,7 @@ use common_query::prelude::GREPTIME_PHYSICAL_TABLE; use common_telemetry::{tracing, warn}; use opentelemetry_proto::tonic::collector::logs::v1::ExportLogsServiceRequest; use opentelemetry_proto::tonic::collector::trace::v1::ExportTraceServiceRequest; +use operator::insert::{admit_row_insert_batches, admit_write}; use otel_arrow_rust::proto::opentelemetry::collector::metrics::v1::ExportMetricsServiceRequest; use pipeline::{GreptimePipelineParams, PipelineWay}; use servers::error::{self, AuthSnafu, Result as ServerResult}; @@ -141,6 +142,10 @@ impl OpenTelemetryProtocolHandler for Instance { self.check_row_insert_permission(&requests, &ctx, PermissionReq::Action(OTLP_WRITE)) .context(AuthSnafu)?; + let ctx = admit_write(rows as u64, &ctx) + .await + .map_err(BoxedError::new) + .context(error::ExecuteGrpcQuerySnafu)?; self.cache_otlp_legacy(&input_names, &ctx, is_legacy)?; OTLP_METRICS_ROWS.inc_by(rows as u64); @@ -239,6 +244,12 @@ impl OpenTelemetryProtocolHandler for Instance { let targets = trace_permission_targets(&table_name, &spans, &ctx); self.check_table_permission(&ctx, PermissionReq::Action(OTLP_WRITE), targets) .context(AuthSnafu)?; + // Count the external spans once, before chunking, retries, or derived tables. + let rows = spans.iter().map(|group| group.spans.len() as u64).sum(); + let ctx = admit_write(rows, &ctx) + .await + .map_err(BoxedError::new) + .context(error::ExecuteGrpcQuerySnafu)?; self.ingest_trace_spans(&pipeline, table_name, spans, &conventions, ctx) .await } @@ -284,11 +295,15 @@ impl OpenTelemetryProtocolHandler for Instance { ) .await?; - let batches = opt_req.as_req_iter(ctx).collect::>(); + let mut batches = opt_req.as_req_iter(ctx).collect::>(); for (temp_ctx, requests) in &batches { self.check_row_insert_permission(requests, temp_ctx, PermissionReq::Action(OTLP_WRITE)) .context(AuthSnafu)?; } + admit_row_insert_batches(&mut batches) + .await + .map_err(BoxedError::new) + .context(error::ExecuteGrpcQuerySnafu)?; let mut outputs = Vec::with_capacity(batches.len()); for (temp_ctx, requests) in batches { diff --git a/src/frontend/src/instance/otlp/README.md b/src/frontend/src/instance/otlp/README.md index 0ce3e94a26a..c123f1caae1 100644 --- a/src/frontend/src/instance/otlp/README.md +++ b/src/frontend/src/instance/otlp/README.md @@ -68,6 +68,10 @@ The parser in [`servers/src/otlp/trace/span.rs`](../../../../servers/src/otlp/trace/span.rs) produces one `TraceSpanGroup` per resource/scope pair. +After permission checks, admission counts all spans in the request once per +database. Chunks, schema retries, and derived lookup-table writes retain that +admission; their actual writes still contribute WCU accounting. + `ingest_trace_spans` in [`trace_ingest.rs`](trace_ingest.rs) splits the spans in each group into owned chunks. `trace_ingest_chunk_size` defaults to 512; setting it to 0 disables splitting. The option is defined in diff --git a/src/frontend/src/instance/prom_store.rs b/src/frontend/src/instance/prom_store.rs index 9aa228715c0..d972abc9112 100644 --- a/src/frontend/src/instance/prom_store.rs +++ b/src/frontend/src/instance/prom_store.rs @@ -527,11 +527,16 @@ impl PromStoreProtocolHandler for Instance { ) -> ServerResult>> { let mut prepared = Vec::with_capacity(requests.len()); for (ctx, request) in requests { - prepared.push(self.prepare_prom_store_write(request, ctx).await?); + let (request, ctx) = self.prepare_prom_store_write(request, ctx).await?; + prepared.push((ctx, request)); } + operator::insert::admit_row_insert_batches(&mut prepared) + .await + .map_err(BoxedError::new) + .context(error::ExecuteGrpcQuerySnafu)?; let mut outputs = Vec::with_capacity(prepared.len()); - for (request, ctx) in prepared { + for (ctx, request) in prepared { let output = self.write_prepared(request, ctx, with_metric_engine).await; let failed = output.is_err(); outputs.push(output); @@ -765,9 +770,13 @@ impl PromStoreProtocolHandler for ExportMetricHandler { async fn write_all( &self, - requests: Vec<(QueryContextRef, RowInsertRequests)>, + mut requests: Vec<(QueryContextRef, RowInsertRequests)>, with_metric_engine: bool, ) -> ServerResult>> { + operator::insert::admit_row_insert_batches(&mut requests) + .await + .map_err(BoxedError::new) + .context(error::ExecuteGrpcQuerySnafu)?; let mut outputs = Vec::with_capacity(requests.len()); for (ctx, request) in requests { let output = self.write_prepared(request, ctx, with_metric_engine).await; diff --git a/src/frontend/src/server.rs b/src/frontend/src/server.rs index 532b0d08b11..017bba53fba 100644 --- a/src/frontend/src/server.rs +++ b/src/frontend/src/server.rs @@ -45,6 +45,7 @@ use servers::postgres::PostgresServer; use servers::request_memory_limiter::ServerMemoryLimiter; use servers::server::{Server, ServerHandlers}; use servers::tls::{ReloadableTlsServerConfig, maybe_watch_server_tls_config}; +use session::context::Channel; use snafu::ResultExt; use tonic::Status; @@ -280,11 +281,15 @@ where .flight_handler(flight_handler) .add_layer(axum::middleware::from_fn_with_state( self.instance.clone(), - async move |State(state): State>, request: Request, next: Next| { + move |State(state): State>, mut request: Request, next: Next| async move { if state.is_suspended() { let status = Status::from(servers::error::SuspendedSnafu.build()); return status.into_http(); } + // The listener owns this marker; clients cannot set request extensions. + if !external { + request.extensions_mut().insert(Channel::Internal); + } next.run(request).await }, )); @@ -612,6 +617,10 @@ mod tests { &self, request: Request, ) -> std::result::Result>, Status> { + assert_eq!( + request.extensions().get::(), + Some(&Channel::Internal) + ); self.do_get_calls.fetch_add(1, Ordering::SeqCst); self.inner.do_get(request).await } @@ -620,6 +629,10 @@ mod tests { &self, request: Request>, ) -> std::result::Result>, Status> { + assert_eq!( + request.extensions().get::(), + Some(&Channel::Internal) + ); self.do_put_calls.fetch_add(1, Ordering::SeqCst); self.inner.do_put(request).await } diff --git a/src/operator/AGENTS.md b/src/operator/AGENTS.md index 761d8e5080a..62853a18f61 100644 --- a/src/operator/AGENTS.md +++ b/src/operator/AGENTS.md @@ -38,3 +38,11 @@ cargo sqlness bare -t ``` Check standalone and distributed routing when changing region dispatch. + +Write admission uses fallible `write_meter!` calls in `Inserter::do_request` +and `handle_bulk_insert`, before flow side effects and datanode dispatch. Bulk +records carry row counts and the request channel with zero WCU. Validate both +the default `meter-macros/noop` build and active metering when changing this path. +Finite requests split internally use `admit_write` / `admit_row_insert_batches` +before dispatch. Their database-scoped `QueryContext` admission prevents a second +row debit while `do_request` still records WCU for the actual writes. diff --git a/src/operator/src/bulk_insert.rs b/src/operator/src/bulk_insert.rs index 51a08d2f324..f1f57a64479 100644 --- a/src/operator/src/bulk_insert.rs +++ b/src/operator/src/bulk_insert.rs @@ -30,7 +30,9 @@ use common_grpc::flight::{FlightEncoder, FlightMessage, record_batch_to_ipc}; use common_telemetry::error; use common_telemetry::tracing_context::TracingContext; use futures::future::{join_all, try_join_all}; -use session::context::QueryContextRef; +use meter_core::data::MeterRecord; +use meter_macros::write_meter; +use session::context::{Channel, QueryContextRef}; use snafu::{ResultExt, ensure}; use store_api::storage::RegionId; use table::TableRef; @@ -128,6 +130,7 @@ impl Inserter { record_batch: RecordBatch, schema_bytes: Bytes, skip_wal: bool, + channel: Channel, ) -> Result { let table_info = table.table_info(); let table_id = table_info.table_id(); @@ -137,6 +140,18 @@ impl Inserter { return Ok(0); } + // The zero value is WCU, not bytes. Bulk writes have no WCU accounting; + // preserve that behavior while admitting their rows before dispatch. + write_meter!(MeterRecord::new( + table_info.catalog_name.clone(), + table_info.schema_name.clone(), + 0, + record_batch.num_rows() as u64, + channel as u8, + )) + .await + .context(error::WriteRejectedSnafu)?; + let body_size = raw_flight_data.data_body.len(); // TODO(yingwen): Fill record batch impure default values. Note that we should override `raw_flight_data` if we have to fill defaults. // notify flownode to update dirty timestamps if flow is configured. diff --git a/src/operator/src/error.rs b/src/operator/src/error.rs index 5a005250b94..40f968bc413 100644 --- a/src/operator/src/error.rs +++ b/src/operator/src/error.rs @@ -122,6 +122,14 @@ pub enum Error { source: BoxedError, }, + #[snafu(display("Write rejected: {error}"))] + WriteRejected { + #[snafu(source)] + error: meter_core::collect::WriteRejected, + #[snafu(implicit)] + location: Location, + }, + #[snafu(display("Failed to insert data"))] RequestInserts { #[snafu(implicit)] @@ -1066,6 +1074,7 @@ impl ErrorExt for Error { Error::AlterExprToRequest { source, .. } => source.status_code(), Error::External { source, .. } => source.status_code(), Error::BatchFlush { source, .. } => source.status_code(), + Error::WriteRejected { .. } => StatusCode::RateLimited, Error::FindTablePartitionRule { source, .. } | Error::SplitInsert { source, .. } | Error::SplitDelete { source, .. } @@ -1122,7 +1131,7 @@ impl ErrorExt for Error { fn retry_hint(&self) -> RetryHint { match self { Error::ReadObject { error, .. } => retry_hint_from_opendal_error(error), - Error::ReadParquetMetadata { .. } => RetryHint::Retryable, + Error::ReadParquetMetadata { .. } | Error::WriteRejected { .. } => RetryHint::Retryable, Error::InvalidateTableCache { source, .. } | Error::ExecuteDdl { source, .. } | Error::RequestInserts { source, .. } diff --git a/src/operator/src/insert.rs b/src/operator/src/insert.rs index 35727ae9b35..989a4aa399d 100644 --- a/src/operator/src/insert.rs +++ b/src/operator/src/insert.rs @@ -48,6 +48,7 @@ use common_telemetry::tracing_context::TracingContext; use common_telemetry::{debug, error, warn}; use datatypes::schema::SkippingIndexOptions; use futures_util::future; +use meter_core::data::MeterRecord; use meter_macros::write_meter; use partition::manager::PartitionRuleManagerRef; use session::context::QueryContextRef; @@ -78,6 +79,7 @@ use crate::batcher::PendingRowsBatcher; use crate::error::{ CatalogSnafu, ColumnOptionsSnafu, CreatePartitionRulesSnafu, FindRegionLeaderSnafu, InvalidInsertRequestSnafu, JoinTaskSnafu, RequestInsertsSnafu, Result, TableNotFoundSnafu, + WriteRejectedSnafu, }; use crate::expr_helper; use crate::region_req_factory::RegionRequestFactory; @@ -386,12 +388,19 @@ impl Inserter { }, instant_requests: RegionInsertRequests::default(), }; + let table_info = table_infos.values().next(); + let catalog = table_info.map_or(ctx.current_catalog(), |info| info.catalog_name.as_str()); + let schema = + table_info.map_or_else(|| ctx.current_schema(), |info| info.schema_name.clone()); let write_cost = write_meter!( - ctx.current_catalog(), - ctx.current_schema(), + catalog, + &schema, metered, + ctx.write_rows_to_admit(catalog, &schema, count_insert_rows(&metered)?), ctx.channel() as u8 - ); + ) + .await + .context(WriteRejectedSnafu)?; prepared.retain(|(_, batch)| batch.num_rows() != 0); let results = if prepared.is_empty() { Vec::new() @@ -399,8 +408,7 @@ impl Inserter { // One original request shares admission across all table submissions. let permit = batcher.acquire().await?; let submissions = prepared.into_iter().map(|(info, batch)| { - // Routing uses the target database; metering above retains the - // original request context, including fully qualified SQL writes. + // Route to the same target database used for admission above. let mut target_ctx = ctx.fork(); target_ctx.set_current_catalog(&info.catalog_name); target_ctx.set_current_schema(&info.schema_name); @@ -542,6 +550,72 @@ impl Inserter { } } +/// Admits a finite request before it is split into internal writes. +/// The returned context preserves accounting while preventing a second row debit. +pub async fn admit_write(rows: u64, ctx: &QueryContextRef) -> Result { + // The zero value is WCU: this record only admits rows. Actual inserts retain + // their existing WCU accounting, so charging here would count it twice. + write_meter!(MeterRecord::new( + ctx.current_catalog().to_string(), + ctx.current_schema(), + 0, + ctx.write_rows_to_admit(ctx.current_catalog(), &ctx.current_schema(), rows), + ctx.channel() as u8, + )) + .await + .context(WriteRejectedSnafu)?; + Ok(Arc::new(ctx.with_write_admission())) +} + +/// Admits all database totals before dispatching any batch of a finite request. +/// Each batch keeps its own protocol options and target database. +pub async fn admit_row_insert_batches( + batches: &mut [(QueryContextRef, RowInsertRequests)], +) -> Result<()> { + let mut totals = BTreeMap::<_, (QueryContextRef, u64)>::new(); + for (ctx, requests) in batches.iter() { + let catalog = ctx.current_catalog(); + let schema = ctx.current_schema(); + if ctx.write_rows_to_admit(catalog, &schema, 1) == 0 { + continue; + } + let (_, total) = totals + .entry((catalog.to_string(), schema.clone())) + .or_insert_with(|| (ctx.clone(), 0)); + for rows in requests.inserts.iter().filter_map(|r| r.rows.as_ref()) { + *total = + total + .checked_add(rows.rows.len() as u64) + .context(InvalidInsertRequestSnafu { + reason: "Insert row count exceeds u64::MAX", + })?; + } + } + for (ctx, rows) in totals.values() { + admit_write(*rows, ctx).await?; + } + for (ctx, _) in batches { + *ctx = Arc::new(ctx.with_write_admission()); + } + Ok(()) +} + +fn count_insert_rows(requests: &InstantAndNormalInsertRequests) -> Result { + requests + .normal_requests + .requests + .iter() + .chain(&requests.instant_requests.requests) + .filter_map(|request| request.rows.as_ref()) + .try_fold(0u64, |total, rows| { + total + .checked_add(rows.rows.len() as u64) + .context(InvalidInsertRequestSnafu { + reason: "Insert row count exceeds u64::MAX", + }) + }) +} + impl Inserter { async fn do_request( &self, @@ -552,12 +626,21 @@ impl Inserter { // Fill impure default values in the request let requests = fill_reqs_with_impure_default(table_infos, requests)?; + // All tables in a batch resolve to the same database. Qualified SQL + // inserts may target a different database than the session's current one. + let table_info = table_infos.values().next(); + let catalog = table_info.map_or(ctx.current_catalog(), |info| info.catalog_name.as_str()); + let schema = + table_info.map_or_else(|| ctx.current_schema(), |info| info.schema_name.clone()); let write_cost = write_meter!( - ctx.current_catalog(), - ctx.current_schema(), + catalog, + schema.clone(), requests, + ctx.write_rows_to_admit(catalog, &schema, count_insert_rows(&requests)?), ctx.channel() as u8 - ); + ) + .await + .context(WriteRejectedSnafu)?; let request_factory = RegionRequestFactory::new(RegionRequestHeader { tracing_context: TracingContext::from_current_span().to_w3c(), dbname: ctx.get_db_string(), @@ -1721,7 +1804,9 @@ mod tests { use table::metadata::{TableInfoBuilder, TableMetaBuilder, TableType}; use crate::insert::*; - use crate::test_util::{create_partition_rule_manager, prepare_mocked_backend}; + use crate::test_util::{ + create_partition_rule_manager, new_test_table_info, prepare_mocked_backend, + }; fn make_table_ref_with_schema( ts_name: &str, @@ -1909,6 +1994,340 @@ mod tests { assert!(table_is_native_histogram(&table)); } + // Keep global meter registration in one test, isolated by nextest's per-test process. + #[tokio::test] + async fn test_write_meter_admission() { + use std::cell::Cell; + use std::sync::Mutex; + use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; + + use api::region::RegionResponse; + use api::v1::region::region_request::Body; + use arrow::array::{Int32Array, TimestampMillisecondArray}; + use arrow::record_batch::RecordBatch; + use bytes::Bytes; + use common_error::ext::{ErrorExt, RetryHint}; + use common_error::status_code::StatusCode; + use common_grpc::flight::{FlightEncoder, FlightMessage}; + use common_meta::ddl::test_util::datanode_handler::DatanodeWatcher; + use futures::future::BoxFuture; + use meter_core::ItemCalculator; + use meter_core::collect::{Collect, WriteRejected}; + use meter_core::data::MeterRecord; + use meter_core::global::global_registry; + use session::context::Channel; + + const CATALOG: &str = "write_meter_test"; + + #[derive(Default)] + struct Meter { + reject: AtomicBool, + attempts: Mutex>, + accepted_value: AtomicU64, + } + + impl Collect for Meter { + fn on_write( + &self, + record: MeterRecord, + ) -> BoxFuture<'_, std::result::Result<(), WriteRejected>> { + Box::pin(async move { + if record.catalog != CATALOG { + return Ok(()); + } + let value = record.value; + self.attempts.lock().unwrap().push(record); + if self.reject.load(Ordering::Relaxed) { + return Err(WriteRejected::new("database row quota exhausted")); + } + self.accepted_value.fetch_add(value, Ordering::Relaxed); + Ok(()) + }) + } + + fn on_read(&self, _: MeterRecord) {} + } + + impl ItemCalculator for Meter { + fn calc(&self, _: &InstantAndNormalInsertRequests) -> u64 { + 17 + } + } + + let kv_backend = prepare_mocked_backend().await; + let partition_manager = create_partition_rule_manager(kv_backend.clone()).await; + let (sender, mut dispatched) = tokio::sync::mpsc::channel(16); + let watcher = DatanodeWatcher::new(sender).with_handler(|_, request| { + let rows = match request.body.unwrap() { + Body::Inserts(requests) => requests + .requests + .iter() + .filter_map(|request| request.rows.as_ref()) + .map(|rows| rows.rows.len()) + .sum(), + // The bulk batches below each contain two rows. + Body::BulkInsert(_) => 2, + body => panic!("unexpected request: {body:?}"), + }; + Ok(RegionResponse::new(rows)) + }); + let flow_cache = Cache::new(100); + let inserter = Inserter::new( + catalog::memory::MemoryCatalogManager::new(), + partition_manager, + Arc::new(MockDatanodeManager::new(watcher)), + Arc::new(new_table_flownode_set_cache( + String::new(), + flow_cache.clone(), + kv_backend, + )), + true, + ); + let mut table_info = new_test_table_info(1, "table_1", [1].into_iter()); + table_info.catalog_name = CATALOG.to_string(); + table_info.schema_name = "target_db".to_string(); + let table_info = Arc::new(table_info); + let table_infos = HashMap::from_iter([(1, table_info.clone())]); + let ctx = Arc::new(QueryContext::with_channel( + DEFAULT_CATALOG_NAME, + DEFAULT_SCHEMA_NAME, + Channel::Postgres, + )); + let rows_request = || { + let build = |num_rows| RegionInsertRequests { + requests: vec![RegionInsertRequest { + region_id: RegionId::new(1, 1).as_u64(), + rows: Some(Rows { + schema: vec![], + rows: vec![api::v1::Row { values: vec![] }; num_rows], + }), + ..Default::default() + }], + }; + InstantAndNormalInsertRequests { + normal_requests: build(3), + instant_requests: build(2), + } + }; + + // No collector or calculator: ordinary OSS insertion still succeeds. + let output = inserter + .do_request(rows_request(), &table_infos, &ctx) + .await + .unwrap(); + assert_eq!(output.meta.cost, 0); + assert!(matches!(output.data, OutputData::AffectedRows(3))); + dispatched.try_recv().unwrap(); + + let meter = Arc::new(Meter::default()); + global_registry().set_collector(meter.clone()); + global_registry().register_calculator(meter.clone()); + // The dependency controls noop mode; exercise both builds with this test. + let enabled = Cell::new(false); + write_meter!({ + enabled.set(true); + MeterRecord::new("probe".into(), "probe".into(), 0, 0, 0) + }) + .await + .unwrap(); + let enabled = enabled.get(); + + let output = inserter + .do_request(rows_request(), &table_infos, &ctx) + .await + .unwrap(); + assert_eq!(output.meta.cost, if enabled { 17 } else { 0 }); + assert!(matches!(output.data, OutputData::AffectedRows(3))); + dispatched.try_recv().unwrap(); + assert!(flow_cache.contains_key(&1)); + flow_cache.invalidate_all(); + meter.reject.store(true, Ordering::Relaxed); + let result = inserter + .do_request(rows_request(), &table_infos, &ctx) + .await; + if enabled { + let error = result.unwrap_err(); + assert_eq!(error.status_code(), StatusCode::RateLimited); + assert_eq!(error.retry_hint(), RetryHint::Retryable); + assert!(error.to_string().contains("database row quota exhausted")); + assert!(dispatched.try_recv().is_err()); + assert!( + !flow_cache.contains_key(&1), + "rejected write reached flow mirroring" + ); + assert_eq!(meter.accepted_value.load(Ordering::Relaxed), 17); + let attempts = meter.attempts.lock().unwrap(); + assert_eq!(attempts.len(), 2); + for record in attempts.iter() { + assert_eq!(record.catalog, CATALOG); + assert_eq!(record.schema, "target_db"); + assert_eq!( + (record.rows, record.value, record.source), + (5, 17, Channel::Postgres as u8) + ); + } + } else { + assert_eq!(result.unwrap().meta.cost, 0); + dispatched.try_recv().unwrap(); + assert!(meter.attempts.lock().unwrap().is_empty()); + } + meter.attempts.lock().unwrap().clear(); + + // Rejection must also precede admission to the table batcher's queue. + if enabled { + let batcher: Arc = Arc::new(UnexpectedBatcher); + let rows = Rows { + schema: vec![ + api::v1::helper::tag_column_schema("a", ColumnDataType::Int32), + time_index_column_schema("ts", ColumnDataType::TimestampMillisecond), + field_column_schema("b", ColumnDataType::Int32), + ], + rows: vec![api::v1::Row { + values: vec![ + api::v1::value::ValueData::I32Value(60).into(), + Value { + value_data: Some(api::v1::value::ValueData::TimestampMillisecondValue( + 0, + )), + }, + api::v1::value::ValueData::I32Value(0).into(), + ], + }], + }; + let error = inserter + .submit_table_rows(rows, table_info.clone(), ctx.clone(), &batcher) + .await + .unwrap_err(); + assert_eq!(error.status_code(), StatusCode::RateLimited); + let mut attempts = meter.attempts.lock().unwrap(); + assert_eq!(attempts.len(), 1); + assert_eq!(attempts[0].catalog, CATALOG); + assert_eq!(attempts[0].schema, "target_db"); + assert_eq!(attempts[0].rows, 1); + attempts.clear(); + } + + let table = Arc::new(table::Table::new( + table_info.clone(), + table::metadata::FilterPushDownType::Unsupported, + Arc::new(DummyDataSource), + )); + let batch = RecordBatch::try_new( + table_info.meta.schema.arrow_schema().clone(), + vec![ + Arc::new(Int32Array::from(vec![60, 70])), + Arc::new(TimestampMillisecondArray::from(vec![0, 1])), + Arc::new(Int32Array::from(vec![0, 0])), + ], + ) + .unwrap(); + let bulk_insert = |batch: RecordBatch| { + let flight_data = FlightEncoder::default() + .encode(FlightMessage::RecordBatch(batch.clone())) + .into_iter() + .next() + .unwrap(); + inserter.handle_bulk_insert( + table.clone(), + flight_data, + batch, + Bytes::new(), + false, + Channel::Grpc, + ) + }; + // Empty batches bypass admission even while the collector rejects. + assert_eq!(bulk_insert(batch.slice(0, 0)).await.unwrap(), 0); + assert!(meter.attempts.lock().unwrap().is_empty()); + assert!(dispatched.try_recv().is_err()); + + // A collector change between batches takes effect on the very next batch. + for reject in [false, true, false] { + meter.reject.store(reject, Ordering::Relaxed); + let result = bulk_insert(batch.clone()).await; + if enabled && reject { + let error = result.unwrap_err(); + assert_eq!(error.status_code(), StatusCode::RateLimited); + assert_eq!(error.retry_hint(), RetryHint::Retryable); + assert!(dispatched.try_recv().is_err()); + } else { + assert_eq!(result.unwrap(), 2); + dispatched.try_recv().unwrap(); + } + } + { + let attempts = meter.attempts.lock().unwrap(); + assert_eq!(attempts.len(), if enabled { 3 } else { 0 }); + for record in attempts.iter() { + assert_eq!(record.catalog, CATALOG); + assert_eq!(record.schema, "target_db"); + assert_eq!( + (record.rows, record.value, record.source), + (2, 0, Channel::Grpc as u8) + ); + } + } + assert_eq!( + meter.accepted_value.load(Ordering::Relaxed), + if enabled { 17 } else { 0 } + ); + + // Aggregate repeated database targets and retain admission across nested + // batching without changing the caller's reusable context. + meter.attempts.lock().unwrap().clear(); + let original = Arc::new(QueryContext::with_channel( + CATALOG, + "a", + Channel::Prometheus, + )); + let mut batches = ["a", "b", "a"].map(|schema| { + let ctx = if schema == "a" { + original.clone() + } else { + Arc::new(QueryContext::with_channel( + CATALOG, + schema, + Channel::Prometheus, + )) + }; + ( + ctx, + RowInsertRequests { + inserts: vec![RowInsertRequest { + table_name: "data".into(), + rows: Some(Rows { + schema: vec![], + rows: vec![api::v1::Row::default(); 2], + }), + }], + }, + ) + }); + admit_row_insert_batches(&mut batches).await.unwrap(); + admit_row_insert_batches(&mut batches).await.unwrap(); + assert_eq!(original.write_rows_to_admit(CATALOG, "a", 4), 4); + for (ctx, _) in &batches { + assert_eq!( + ctx.write_rows_to_admit(CATALOG, &ctx.current_schema(), 2), + 0 + ); + assert_eq!(ctx.channel(), Channel::Prometheus); + } + let attempts = meter.attempts.lock().unwrap(); + let totals = attempts + .iter() + .map(|r| (r.schema.as_str(), r.rows, r.value)) + .collect::>(); + assert_eq!( + totals, + if enabled { + vec![("a", 4, 0), ("b", 2, 0)] + } else { + vec![] + } + ); + } + #[test] fn test_skip_wal_does_not_change_table_options() { check_skip_wal_does_not_change_table_options(false); diff --git a/src/operator/src/statement/ddl.rs b/src/operator/src/statement/ddl.rs index 2f020c28e68..0bebd3c1295 100644 --- a/src/operator/src/statement/ddl.rs +++ b/src/operator/src/statement/ddl.rs @@ -101,7 +101,7 @@ use table::dist_table::DistTable; use table::metadata::{self, TableId, TableInfo, TableMeta, TableType}; use table::requests::{ AlterKind, AlterTableRequest, AnnotationContext, COMMENT_KEY, DDL_TIMEOUT, DDL_WAIT, - TableOptions, validate_and_normalize_annotation_options, + INGEST_ROWS_RATE_LIMIT_KEY, TableOptions, validate_and_normalize_annotation_options, }; use table::table_name::TableName; use table::table_reference::TableReference; @@ -357,11 +357,10 @@ impl StatementExecutor { .map(|v| v.into_inner()); let create_expr = &mut expr_helper::create_to_expr(&stmt, &ctx)?; - // Don't inherit schema-level TTL/compaction options into table options: - // TTL is applied during compaction, and `compaction.*` is handled separately. + // TTL and compaction options are handled separately; ingestion quota is database-only. if let Some(schema_options) = schema_options { for (key, value) in schema_options.extra_options.iter() { - if key.starts_with("compaction.") { + if key == INGEST_ROWS_RATE_LIMIT_KEY || key.starts_with("compaction.") { continue; } create_expr diff --git a/src/servers/AGENTS.md b/src/servers/AGENTS.md index 2a446033f62..006ac59771c 100644 --- a/src/servers/AGENTS.md +++ b/src/servers/AGENTS.md @@ -25,6 +25,12 @@ provided through handler traits, mostly implemented by `frontend`. after bumping it. - Auth or `QueryContext` changes must cover HTTP, gRPC, MySQL, and PostgreSQL entry points. +- Internal frontend gRPC listeners set `Channel::Internal` in server-owned request + extensions. Handlers preserve it through query execution and metering; client + flow metadata does not grant an internal channel. +- Finite requests split by wire handlers (Prometheus batching and OpenTSDB + summary/details) admit all rows before dispatch. Admission is process-local + `QueryContext` state; never accept it from headers or serialized contexts. - New protocols or externally visible routes also require frontend service and configuration wiring. diff --git a/src/servers/Cargo.toml b/src/servers/Cargo.toml index 8ffd97e2c0f..0e2d896db8c 100644 --- a/src/servers/Cargo.toml +++ b/src/servers/Cargo.toml @@ -82,6 +82,8 @@ jsonb.workspace = true lazy_static.workspace = true log-query.workspace = true loki-proto.workspace = true +meter-core.workspace = true +meter-macros.workspace = true metric-engine.workspace = true mime_guess = "2.0" mysql_common = "0.38" diff --git a/src/servers/src/batcher/logical_table.rs b/src/servers/src/batcher/logical_table.rs index de872ac326b..9543b0d3a88 100644 --- a/src/servers/src/batcher/logical_table.rs +++ b/src/servers/src/batcher/logical_table.rs @@ -40,6 +40,8 @@ use common_query::prelude::GREPTIME_PHYSICAL_TABLE; use common_runtime::spawn_global; use common_telemetry::{debug, error, warn}; use datatypes::timestamp::append_timestamps; +use meter_core::data::MeterRecord; +use meter_macros::write_meter; use partition::manager::PartitionRuleManagerRef; use session::context::QueryContextRef; use snafu::{OptionExt, ResultExt}; @@ -276,6 +278,21 @@ impl LogicalTablePendingRowsBatcher { return Ok(0); } + // Flushes dispatch directly to datanodes, so admit once before enqueueing. + write_meter!(MeterRecord::new( + ctx.current_catalog().to_string(), + ctx.current_schema(), + 0, + ctx.write_rows_to_admit( + ctx.current_catalog(), + &ctx.current_schema(), + total_rows as u64 + ), + ctx.channel() as u8, + )) + .await + .context(error::WriteRejectedSnafu)?; + let permit = { let _timer = PENDING_ROWS_BATCH_INGEST_STAGE_ELAPSED .with_label_values(&["submit_acquire_inflight_permit"]) diff --git a/src/servers/src/error.rs b/src/servers/src/error.rs index 30997a8536a..47f9953a187 100644 --- a/src/servers/src/error.rs +++ b/src/servers/src/error.rs @@ -62,6 +62,14 @@ pub enum Error { #[snafu(display("Pending rows batcher channel closed"))] BatcherChannelClosed, + #[snafu(display("Write rejected: {error}"))] + WriteRejected { + #[snafu(source)] + error: meter_core::collect::WriteRejected, + #[snafu(implicit)] + location: Location, + }, + #[snafu(display("Unsupported data type: {}, reason: {}", data_type, reason))] UnsupportedDataType { data_type: ConcreteDataType, @@ -855,7 +863,7 @@ impl ErrorExt for Error { Suspended { .. } => StatusCode::Suspended, - MemoryLimitExceeded { .. } => StatusCode::RateLimited, + MemoryLimitExceeded { .. } | WriteRejected { .. } => StatusCode::RateLimited, GreptimeProto { source, .. } => source.status_code(), Partition { source, .. } => source.status_code(), @@ -901,7 +909,7 @@ impl ErrorExt for Error { MemoryLimitExceeded { source, .. } => source.retry_hint(), CollectRecordbatch { source, .. } => source.retry_hint(), - TooManyConcurrentRequests { .. } => RetryHint::Retryable, + TooManyConcurrentRequests { .. } | WriteRejected { .. } => RetryHint::Retryable, _ => RetryHint::NonRetryable, } diff --git a/src/servers/src/grpc/context_auth.rs b/src/servers/src/grpc/context_auth.rs index 2e0bda0939c..bb6f2d27d2a 100644 --- a/src/servers/src/grpc/context_auth.rs +++ b/src/servers/src/grpc/context_auth.rs @@ -34,9 +34,10 @@ use crate::http::AUTHORIZATION_HEADER; use crate::http::header::constants::GREPTIME_DB_HEADER_NAME; use crate::metrics::METRIC_AUTH_FAILURE; -/// Create a query context from the grpc metadata. +/// Create a query context from gRPC metadata and server-owned request extensions. pub fn create_query_context_from_grpc_metadata( headers: &MetadataMap, + extensions: &http::Extensions, ) -> TonicResult { let (catalog, schema) = if let Some(db) = extract_header(headers, &[GREPTIME_DB_HEADER_NAME])? { parse_catalog_and_schema_from_db_string(db) @@ -50,7 +51,12 @@ pub fn create_query_context_from_grpc_metadata( let ctx = QueryContextBuilder::default() .current_catalog(catalog) .current_schema(schema) - .channel(Channel::Grpc) + .channel( + extensions + .get::() + .copied() + .unwrap_or(Channel::Grpc), + ) .build(); // OTEL Arrow uses ordinary inserts. Accept only its request-level WAL hint, // leaving unrelated hints and reserved internal extensions unchanged. @@ -183,11 +189,28 @@ mod tests { use super::*; + #[test] + fn test_channel_comes_from_server_extensions() { + let mut headers = MetadataMap::new(); + headers.insert( + "x-greptime-flow-extensions", + r#"[["flow.return_region_seq","true"]]"#.parse().unwrap(), + ); + headers.insert(HINTS_KEY, "channel=internal".parse().unwrap()); + let mut extensions = http::Extensions::new(); + let ctx = create_query_context_from_grpc_metadata(&headers, &extensions).unwrap(); + assert_eq!(ctx.channel(), Channel::Grpc); + + extensions.insert(Channel::Internal); + let ctx = create_query_context_from_grpc_metadata(&headers, &extensions).unwrap(); + assert_eq!(ctx.channel(), Channel::Internal); + } + #[test] fn test_arrow_insert_hint_does_not_accept_reserved_extensions() { let mut headers = MetadataMap::new(); assert_eq!( - create_query_context_from_grpc_metadata(&headers) + create_query_context_from_grpc_metadata(&headers, &Default::default()) .unwrap() .extension(INSERT_SKIP_WAL_HINT), None @@ -198,7 +221,8 @@ mod tests { hints.push_str(&format!(",{key}=external")); } headers.insert(HINTS_KEY, hints.parse().unwrap()); - let ctx = create_query_context_from_grpc_metadata(&headers).unwrap(); + let ctx = + create_query_context_from_grpc_metadata(&headers, &Default::default()).unwrap(); assert_eq!(ctx.skip_wal(), expected); assert_eq!(ctx.extension(INSERT_SKIP_WAL_HINT), None); assert_eq!(ctx.extension("ttl"), None); @@ -219,7 +243,8 @@ mod tests { ("insert_skip_wal=true,insert_skip_wal=invalid", true), ] { headers.insert(HINTS_KEY, hints.parse().unwrap()); - let ctx = create_query_context_from_grpc_metadata(&headers).unwrap(); + let ctx = + create_query_context_from_grpc_metadata(&headers, &Default::default()).unwrap(); assert_eq!(ctx.skip_wal(), expected); } for value in ["", "TRUE", "1", "invalid"] { @@ -227,7 +252,9 @@ mod tests { HINTS_KEY, format!("insert_skip_wal={value}").parse().unwrap(), ); - assert!(create_query_context_from_grpc_metadata(&headers).is_err()); + assert!( + create_query_context_from_grpc_metadata(&headers, &Default::default()).is_err() + ); } } } diff --git a/src/servers/src/grpc/database.rs b/src/servers/src/grpc/database.rs index 79904806aba..2a63a982bef 100644 --- a/src/servers/src/grpc/database.rs +++ b/src/servers/src/grpc/database.rs @@ -21,6 +21,7 @@ use common_query::OutputData; use common_telemetry::{debug, warn}; use futures::StreamExt; use prost::Message; +use session::context::Channel; use tonic::{Request, Response, Status, Streaming}; use crate::grpc::greptime_handler::GreptimeRequestHandler; @@ -46,6 +47,11 @@ impl GreptimeDatabase for DatabaseService { ) -> TonicResult> { let remote_addr = request.remote_addr(); let hints = hint_headers::extract_hints(request.metadata()); + let channel = request + .extensions() + .get::() + .copied() + .unwrap_or(Channel::Grpc); debug!( "GreptimeDatabase::Handle: request from {:?} with hints: {:?}", remote_addr, hints @@ -61,7 +67,7 @@ impl GreptimeDatabase for DatabaseService { let handler = self.handler.clone(); let request_future = async move { let request = request.into_inner(); - let output = handler.handle_request(request, hints).await?; + let output = handler.handle_request(request, hints, channel).await?; let message = match output.data { OutputData::AffectedRows(rows) => GreptimeResponse { header: Some(ResponseHeader { @@ -100,6 +106,11 @@ impl GreptimeDatabase for DatabaseService { ) -> Result, Status> { let remote_addr = request.remote_addr(); let hints = hint_headers::extract_hints(request.metadata()); + let channel = request + .extensions() + .get::() + .copied() + .unwrap_or(Channel::Grpc); debug!( "GreptimeDatabase::HandleRequests: request from {:?} with hints: {:?}", remote_addr, hints @@ -121,7 +132,9 @@ impl GreptimeDatabase for DatabaseService { } else { None }; - let output = handler.handle_request(request, hints.clone()).await?; + let output = handler + .handle_request(request, hints.clone(), channel) + .await?; match output.data { OutputData::AffectedRows(rows) => affected_rows += rows, OutputData::Stream(_) | OutputData::RecordBatches(_) => { diff --git a/src/servers/src/grpc/flight.rs b/src/servers/src/grpc/flight.rs index f0035ed3f8c..874e98b7326 100644 --- a/src/servers/src/grpc/flight.rs +++ b/src/servers/src/grpc/flight.rs @@ -199,12 +199,17 @@ impl FlightCraft for GreptimeRequestHandler { let mut hints = hint_headers::extract_hints(request.metadata()); hints.extend(extract_flow_extensions(request.metadata())?); let snapshot_seqs = extract_snapshot_seqs(request.metadata())?; + let channel = request + .extensions() + .get::() + .copied() + .unwrap_or(Channel::Grpc); let ticket = request.into_inner().ticket; let request = GreptimeRequest::decode(ticket.as_ref()).context(error::InvalidFlightTicketSnafu)?; let query_ctx = - create_query_context(Channel::Grpc, request.header.as_ref(), hints, snapshot_seqs)?; + create_query_context(channel, request.header.as_ref(), hints, snapshot_seqs)?; // Validate flow hint syntax at the transport boundary before dispatching the request. // This does not authorize or execute anything; `handle_request()` below still performs // the normal frontend handling and auth checks before query execution. @@ -249,7 +254,8 @@ impl FlightCraft for GreptimeRequestHandler { let limiter = extensions.get::().cloned(); - let query_ctx = context_auth::create_query_context_from_grpc_metadata(&headers)?; + let query_ctx = + context_auth::create_query_context_from_grpc_metadata(&headers, &extensions)?; context_auth::check_auth(self.user_provider.clone(), &headers, query_ctx.clone()).await?; const MAX_PENDING_RESPONSES: usize = 32; diff --git a/src/servers/src/grpc/greptime_handler.rs b/src/servers/src/grpc/greptime_handler.rs index dcd8c641a64..ad0a693ac59 100644 --- a/src/servers/src/grpc/greptime_handler.rs +++ b/src/servers/src/grpc/greptime_handler.rs @@ -78,9 +78,10 @@ impl GreptimeRequestHandler { &self, request: GreptimeRequest, hints: Vec<(String, String)>, + channel: Channel, ) -> Result { let header = request.header.as_ref(); - let query_ctx = create_query_context(Channel::Grpc, header, hints, HashMap::new())?; + let query_ctx = create_query_context(channel, header, hints, HashMap::new())?; let query = request.request.context(InvalidQuerySnafu { reason: "Expecting non-empty GreptimeRequest.", })?; diff --git a/src/servers/src/grpc/prom_query_gateway.rs b/src/servers/src/grpc/prom_query_gateway.rs index 92bd5d722e5..abec69604f6 100644 --- a/src/servers/src/grpc/prom_query_gateway.rs +++ b/src/servers/src/grpc/prom_query_gateway.rs @@ -47,6 +47,11 @@ pub struct PrometheusGatewayService { impl PrometheusGateway for PrometheusGatewayService { async fn handle(&self, req: Request) -> TonicResult> { let mut is_range_query = false; + let channel = req + .extensions() + .get::() + .copied() + .unwrap_or(Channel::Promql); let inner = req.into_inner(); let prom_query = match inner.promql.context(InvalidQuerySnafu { reason: "Expecting non-empty PromqlRequest.", @@ -80,12 +85,8 @@ impl PrometheusGateway for PrometheusGatewayService { }; let header = inner.header.as_ref(); - let query_ctx = create_query_context( - Channel::Promql, - header, - Default::default(), - Default::default(), - )?; + let query_ctx = + create_query_context(channel, header, Default::default(), Default::default())?; let user_info = auth(self.user_provider.clone(), header, &query_ctx).await?; query_ctx.set_current_user(user_info); diff --git a/src/servers/src/http/opentsdb.rs b/src/servers/src/http/opentsdb.rs index 51667db2fe7..12937a4d184 100644 --- a/src/servers/src/http/opentsdb.rs +++ b/src/servers/src/http/opentsdb.rs @@ -88,22 +88,22 @@ pub async fn put( .collect::>(); ctx.set_channel(Channel::Opentsdb); - let ctx = Arc::new(ctx); + let mut ctx = Arc::new(ctx); if summary || details { opentsdb_handler .preflight(&data_points, ctx.clone()) .await?; + ctx = operator::insert::admit_write(data_points.len() as u64, &ctx) + .await + .map_err(common_error::ext::BoxedError::new) + .context(error::ExecuteGrpcQuerySnafu)?; } let response = if !summary && !details { - if let Err(e) = opentsdb_handler.exec_batch(data_points, ctx.clone()).await { - // Not debugging purpose, failed fast. - return error::InternalSnafu { - err_msg: e.to_string(), - } - .fail(); - } + opentsdb_handler + .exec_batch(data_points, ctx.clone()) + .await?; (HttpStatusCode::NO_CONTENT, Json(OpentsdbPutResponse::Empty)) } else { let mut response = OpentsdbDebuggingResponse { diff --git a/src/servers/src/http/prom_store.rs b/src/servers/src/http/prom_store.rs index ed1c778f701..17db6abae04 100644 --- a/src/servers/src/http/prom_store.rs +++ b/src/servers/src/http/prom_store.rs @@ -380,12 +380,16 @@ async fn preflight_prometheus_rows( prom_store_handler: &PromStoreProtocolHandlerRef, batches: &mut [PromWriteBatch], ) -> Result<()> { - for (ctx, reqs) in batches { + for (ctx, reqs) in batches.iter_mut() { prom_store_handler.pre_write(reqs, ctx.clone()).await?; // Detach from context clones retained by pre-write hooks so the checked // schema cannot change before this prepared batch is written. *ctx = Arc::new(ctx.fork()); } + operator::insert::admit_row_insert_batches(batches) + .await + .map_err(common_error::ext::BoxedError::new) + .context(error::ExecuteGrpcQuerySnafu)?; Ok(()) } @@ -918,7 +922,8 @@ mod tests { #[async_trait] impl PromWriteBatcher for RecordingPromWriteBatcher { - async fn submit(&self, requests: RowInsertRequests, _ctx: QueryContextRef) -> Result { + async fn submit(&self, requests: RowInsertRequests, ctx: QueryContextRef) -> Result { + assert_eq!(ctx.write_rows_to_admit("greptime", "public", 1), 0); record_write_event(&self.events, "batch", &requests); Ok(prom_write_row_count(&requests)) } @@ -942,9 +947,10 @@ mod tests { async fn write_prepared( &self, request: RowInsertRequests, - _ctx: QueryContextRef, + ctx: QueryContextRef, _with_metric_engine: bool, ) -> Result { + assert_eq!(ctx.write_rows_to_admit("greptime", "public", 1), 0); record_write_event(&self.events, "direct", &request); Ok(Output::new_with_affected_rows(0)) } diff --git a/src/servers/src/otel_arrow.rs b/src/servers/src/otel_arrow.rs index 730020161d3..0cf07c76468 100644 --- a/src/servers/src/otel_arrow.rs +++ b/src/servers/src/otel_arrow.rs @@ -84,9 +84,14 @@ impl ArrowMetricsService for OtelArrowServiceHandler Result, Status> { let (mut sender, receiver) = futures::channel::mpsc::channel(100); - let (headers, _, mut incoming_requests) = request.into_parts(); + let (headers, extensions, mut incoming_requests) = request.into_parts(); - let query_ctx = context_auth::create_query_context_from_grpc_metadata(&headers)?; + // OTEL Arrow is currently used only for external ingestion. Its service bypasses + // the frontend router middleware, so even the internal listener defaults to + // Channel::Grpc and remains rate-limited. Before using this path internally, + // propagate the server-owned Channel::Internal marker to this service. + let query_ctx = + context_auth::create_query_context_from_grpc_metadata(&headers, &extensions)?; context_auth::check_auth(self.user_provider.clone(), &headers, query_ctx.clone()).await?; let query_ctx = { let mut ctx = query_ctx.fork(); diff --git a/src/servers/tests/http/opentsdb_test.rs b/src/servers/tests/http/opentsdb_test.rs index eb2f312f94a..6dc2ec51d52 100644 --- a/src/servers/tests/http/opentsdb_test.rs +++ b/src/servers/tests/http/opentsdb_test.rs @@ -28,6 +28,7 @@ use servers::opentsdb::codec::DataPoint; use servers::query_handler::OpentsdbProtocolHandler; use servers::query_handler::sql::SqlQueryHandler; use session::context::QueryContextRef; +use snafu::IntoError; use sql::statements::statement::Statement; use tokio::sync::mpsc; @@ -43,6 +44,11 @@ impl OpentsdbProtocolHandler for DummyInstance { async fn exec(&self, data_points: Vec, _ctx: QueryContextRef) -> Result { let data_point = data_points.first().unwrap(); + if data_point.metric() == "rate_limited" { + return Err(error::WriteRejectedSnafu.into_error( + meter_core::collect::WriteRejected::new("database row quota exhausted"), + )); + } if data_point.metric() == "should_failed" { return error::InternalSnafu { err_msg: "expected", @@ -156,6 +162,14 @@ async fn test_opentsdb_put() { assert_eq!(result.status(), 500); assert_eq!(result.text().await, "{\"error\":\"Internal error: 1003\"}"); + let result = client + .post("/v1/opentsdb/api/put") + .body(create_data_point("rate_limited")) + .send() + .await; + assert_eq!(result.status(), 429); + assert!(result.text().await.contains("database row quota exhausted")); + let mut metrics = vec![]; while let Ok(s) = rx.try_recv() { metrics.push(s); diff --git a/src/session/src/context.rs b/src/session/src/context.rs index 2833abd2578..0bb3f0f6a31 100644 --- a/src/session/src/context.rs +++ b/src/session/src/context.rs @@ -85,6 +85,9 @@ pub struct QueryContext { /// Track which protocol the query comes from. #[builder(default)] channel: Channel, + /// Process-local admission for one database in this request. Never sent over the wire. + #[builder(setter(skip))] + admitted_write: Option<(String, String)>, /// Process id for managing on-going queries #[builder(default)] process_id: u32, @@ -263,6 +266,28 @@ impl QueryContext { fork } + /// Forks a context for internal writes after admitting the complete request. + /// Call only after write admission succeeds for the current database. + pub fn with_write_admission(&self) -> Self { + let mut ctx = self.fork(); + ctx.admitted_write = Some((ctx.current_catalog().to_string(), ctx.current_schema())); + ctx + } + + /// Returns zero for writes covered by this request's database admission. + /// Usage accounting and the original protocol channel remain unchanged. + pub fn write_rows_to_admit(&self, catalog: &str, schema: &str, rows: u64) -> u64 { + if self + .admitted_write + .as_ref() + .is_some_and(|(c, s)| c == catalog && s == schema) + { + 0 + } else { + rows + } + } + pub fn arc() -> QueryContextRef { Arc::new( QueryContextBuilder::default() @@ -608,6 +633,7 @@ impl QueryContextBuilder { .unwrap_or_else(|| Arc::new(ConfigurationVariables::default())), channel, batching_enabled: self.batching_enabled.unwrap_or_default(), + admitted_write: None, process_id: self.process_id.unwrap_or_default(), conn_info: self.conn_info.unwrap_or_default(), protocol_ctx: self.protocol_ctx.unwrap_or_default(), @@ -788,6 +814,26 @@ mod test { assert!(fork.skip_wal()); } + #[test] + fn test_write_admission_is_local_and_database_scoped() { + let ctx = QueryContext::with_channel("greptime", "public", Channel::Otlp); + let admitted = ctx.with_write_admission(); + assert_eq!(ctx.write_rows_to_admit("greptime", "public", 100), 100); + assert_eq!(admitted.write_rows_to_admit("greptime", "public", 100), 0); + assert_eq!(admitted.write_rows_to_admit("greptime", "other", 100), 100); + assert_eq!(admitted.write_rows_to_admit("other", "public", 100), 100); + assert_eq!(admitted.channel(), Channel::Otlp); + assert_eq!( + admitted + .fork() + .write_rows_to_admit("greptime", "public", 100), + 0 + ); + let wire: api::v1::QueryContext = admitted.into(); + let restored = QueryContext::from(wire); + assert_eq!(restored.write_rows_to_admit("greptime", "public", 100), 100); + } + #[test] fn test_skip_wal_is_not_serialized_in_query_context() { let context = QueryContextBuilder::default().skip_wal(true).build(); @@ -808,7 +854,7 @@ mod test { current_schema: "s1".to_string(), timezone: "UTC".to_string(), extensions: HashMap::from([("flow.return_region_seq".to_string(), "true".to_string())]), - channel: Channel::Grpc as u32, + channel: Channel::Internal as u32, snapshot_seqs: Some(api::v1::SnapshotSequences { snapshot_seqs: HashMap::from([(1, 100)]), sst_min_sequences: HashMap::from([(1, 90)]), diff --git a/src/sql/src/parsers/create_parser.rs b/src/sql/src/parsers/create_parser.rs index 82dcd949ae4..bb73870b45f 100644 --- a/src/sql/src/parsers/create_parser.rs +++ b/src/sql/src/parsers/create_parser.rs @@ -1591,6 +1591,18 @@ mod tests { } _ => unreachable!(), } + + let sql = "CREATE DATABASE prometheus with ('ingest_rows_rate_limit'='1000');"; + let result = + ParserContext::create_with_dialect(sql, &GreptimeDbDialect {}, ParseOptions::default()); + let stmts = result.unwrap(); + match &stmts[0] { + Statement::CreateDatabase(c) => { + assert_eq!(c.name.to_string(), "prometheus"); + assert_eq!(c.options.get("ingest_rows_rate_limit").unwrap(), "1000"); + } + _ => unreachable!(), + } } #[test] diff --git a/src/table/src/requests.rs b/src/table/src/requests.rs index c57dc0b009d..10055df4b45 100644 --- a/src/table/src/requests.rs +++ b/src/table/src/requests.rs @@ -106,6 +106,9 @@ pub const DDL_WAIT: &str = "wait"; pub const VALID_DDL_OPTION_KEYS: [&str; 2] = [DDL_TIMEOUT, DDL_WAIT]; +/// The key of ingest rows rate limit option (rows per second, cluster-wide) in database options. +pub const INGEST_ROWS_RATE_LIMIT_KEY: &str = "ingest_rows_rate_limit"; + // Valid option keys when creating a db. static VALID_DB_OPT_KEYS: Lazy> = Lazy::new(|| { let mut set = HashSet::new(); @@ -129,6 +132,7 @@ static VALID_DB_OPT_KEYS: Lazy> = Lazy::new(|| { set.insert(TWCS_INACTIVE_WINDOW_L1_MERGE_TRIGGER); set.insert(TWCS_MAX_OUTPUT_FILE_SIZE); set.insert(SST_FORMAT_KEY); + set.insert(INGEST_ROWS_RATE_LIMIT_KEY); set }); @@ -142,6 +146,12 @@ pub fn validate_database_option_value( key: &str, value: Option<&str>, ) -> std::result::Result<(), &'static str> { + if key == INGEST_ROWS_RATE_LIMIT_KEY { + return value + .and_then(|value| value.parse::().ok()) + .map(|_| ()) + .ok_or("expected a non-negative integer fitting in u64"); + } let (minimum, constraint) = match key { TWCS_TRIGGER_FILE_NUM | TWCS_ACTIVE_WINDOW_TRIGGER_FILE_NUM @@ -933,6 +943,8 @@ mod tests { assert!(validate_database_option( "compaction.twcs.inactive_window.l1_merge_trigger" )); + assert!(validate_database_option(INGEST_ROWS_RATE_LIMIT_KEY)); + assert!(validate_database_option("ingest_rows_rate_limit")); assert!(!validate_database_option("foo")); } @@ -994,6 +1006,27 @@ mod tests { } } + #[test] + fn test_database_ingest_rate_limit_value_boundaries() { + for invalid in [ + None, + Some(""), + Some("abc"), + Some("1000/s"), + Some("-1"), + Some("1.5"), + Some("18446744073709551616"), + ] { + assert!(validate_database_option_value(INGEST_ROWS_RATE_LIMIT_KEY, invalid).is_err()); + } + let maximum = u64::MAX.to_string(); + for valid in ["0", "1", maximum.as_str()] { + assert!( + validate_database_option_value(INGEST_ROWS_RATE_LIMIT_KEY, Some(valid)).is_ok() + ); + } + } + #[test] fn test_serialize_table_options() { let options = TableOptions { diff --git a/tests-integration/src/tests/instance_test.rs b/tests-integration/src/tests/instance_test.rs index f5e94319bac..dc4d3898860 100644 --- a/tests-integration/src/tests/instance_test.rs +++ b/tests-integration/src/tests/instance_test.rs @@ -325,6 +325,57 @@ PARTITION ON COLUMNS (n) ( check_output_stream(output, expected).await; } +#[apply(both_instances_cases)] +async fn test_database_ingest_rate_limit_not_inherited(instance: Arc) { + let frontend = instance.frontend(); + execute_sql( + &frontend, + "CREATE DATABASE limited WITH ('ingest_rows_rate_limit'='1000', 'skip_wal'='true')", + ) + .await; + let ctx = Arc::new(QueryContext::with(DEFAULT_CATALOG_NAME, "limited")); + + for (name, sql) in [ + ("source", "CREATE TABLE source (ts TIMESTAMP TIME INDEX)"), + ("copy", "CREATE TABLE copy LIKE source"), + ] { + execute_sql_with(&frontend, sql, ctx.clone()).await; + let table = frontend + .catalog_manager() + .table(DEFAULT_CATALOG_NAME, "limited", name, None) + .await + .unwrap() + .unwrap(); + let options = &table.table_info().meta.options; + assert!(!options.extra_options.contains_key("ingest_rows_rate_limit")); + assert!(options.skip_wal); + + let output = + execute_sql_with(&frontend, &format!("SHOW CREATE TABLE {name}"), ctx.clone()).await; + let OutputData::RecordBatches(batches) = output.data else { + unreachable!() + }; + let batch = batches.iter().next().unwrap(); + let ddl = batch + .column_by_name("Create Table") + .unwrap() + .as_string::() + .value(0); + assert!(!ddl.contains("ingest_rows_rate_limit")); + execute_sql_with(&frontend, &format!("DROP TABLE {name}"), ctx.clone()).await; + execute_sql_with(&frontend, ddl, ctx.clone()).await; + } + + let output = execute_sql(&frontend, "SHOW CREATE DATABASE limited").await; + assert!( + output + .data + .pretty_print() + .await + .contains("ingest_rows_rate_limit") + ); +} + #[apply(standalone_instance_case)] async fn test_extra_external_table_options(instance: Arc) { let frontend = instance.frontend();