diff --git a/src/servers/src/grpc/context_auth.rs b/src/servers/src/grpc/context_auth.rs index 0a44c0cb47..2e0bda0939 100644 --- a/src/servers/src/grpc/context_auth.rs +++ b/src/servers/src/grpc/context_auth.rs @@ -54,16 +54,17 @@ pub fn create_query_context_from_grpc_metadata( .build(); // OTEL Arrow uses ordinary inserts. Accept only its request-level WAL hint, // leaving unrelated hints and reserved internal extensions unchanged. - for (key, value) in hint_headers::extract_hints(headers) { - if key == INSERT_SKIP_WAL_HINT { - let skip_wal = value.parse::().map_err(|_| { - InvalidParameterSnafu { - reason: format!("Invalid {key} hint: expected true or false, got {value:?}"), - } - .build() - })?; - ctx.set_skip_wal(skip_wal); - } + if let Some((key, value)) = hint_headers::extract_hints(headers) + .into_iter() + .find(|(key, _)| key == INSERT_SKIP_WAL_HINT) + { + let skip_wal = value.parse::().map_err(|_| { + InvalidParameterSnafu { + reason: format!("Invalid {key} hint: expected true or false, got {value:?}"), + } + .build() + })?; + ctx.set_skip_wal(skip_wal); } Ok(Arc::new(ctx)) } @@ -211,6 +212,16 @@ mod tests { } } } + // Only the first matching hint is parsed and applied. + for (hints, expected) in [ + ("insert_skip_wal=true,insert_skip_wal=false", true), + ("insert_skip_wal=false,insert_skip_wal=true", false), + ("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(); + assert_eq!(ctx.skip_wal(), expected); + } for value in ["", "TRUE", "1", "invalid"] { headers.insert( HINTS_KEY, diff --git a/src/servers/src/grpc/greptime_handler.rs b/src/servers/src/grpc/greptime_handler.rs index 7ae59a024b..dcd8c641a6 100644 --- a/src/servers/src/grpc/greptime_handler.rs +++ b/src/servers/src/grpc/greptime_handler.rs @@ -202,7 +202,7 @@ pub fn get_request_type(request: &GreptimeRequest) -> &'static str { pub(crate) fn create_query_context( channel: Channel, header: Option<&RequestHeader>, - mut extensions: Vec<(String, String)>, + extensions: Vec<(String, String)>, snapshot_seqs: HashMap, ) -> Result { let (catalog, schema) = header @@ -240,39 +240,36 @@ pub(crate) fn create_query_context( .channel(channel) .snapshot_seqs(Arc::new(RwLock::new(snapshot_seqs))); - if let Some(x) = extensions - .iter() - .position(|(k, _)| k == READ_PREFERENCE_HINT) - { - let (k, v) = extensions.swap_remove(x); - let Ok(read_preference) = ReadPreference::from_str(&v) else { - return UnknownHintSnafu { - hint: format!("{k}={v}"), - } - .fail(); - }; - ctx_builder = ctx_builder.read_preference(read_preference); - } - for (key, value) in extensions { - if key == INSERT_SKIP_WAL_HINT { - let skip_wal = value.parse::().map_err(|_| { - UnknownHintSnafu { - hint: format!("{key}={value}"), - } - .build() - })?; - ctx_builder = ctx_builder.skip_wal(skip_wal); - continue; + match key.as_str() { + READ_PREFERENCE_HINT => { + let Ok(read_preference) = ReadPreference::from_str(&value) else { + return UnknownHintSnafu { + hint: format!("{key}={value}"), + } + .fail(); + }; + ctx_builder = ctx_builder.read_preference(read_preference); + } + INSERT_SKIP_WAL_HINT => { + let skip_wal = value.parse::().map_err(|_| { + UnknownHintSnafu { + hint: format!("{key}={value}"), + } + .build() + })?; + ctx_builder = ctx_builder.skip_wal(skip_wal); + } + _ if is_reserved_extension_key(&key) => { + debug!( + key = key.as_str(), + "Ignoring reserved external query context extension key" + ); + } + _ => { + ctx_builder = ctx_builder.set_extension(key, value); + } } - if is_reserved_extension_key(&key) { - debug!( - key = key.as_str(), - "Ignoring reserved external query context extension key" - ); - continue; - } - ctx_builder = ctx_builder.set_extension(key, value); } Ok(ctx_builder.build().into()) } @@ -381,6 +378,32 @@ mod tests { assert_eq!(ctx.extension(INSERT_SKIP_WAL_HINT), None); } + #[test] + fn test_create_query_context_read_preference_duplicates() { + for (values, valid) in [ + (["leader", "LEADER"], true), + (["invalid", "leader"], false), + (["leader", "invalid"], false), + ] { + let result = create_query_context( + Channel::Grpc, + None, + values + .into_iter() + .map(|value| (READ_PREFERENCE_HINT.to_string(), value.to_string())) + .collect(), + HashMap::new(), + ); + if valid { + let ctx = result.unwrap(); + assert!(matches!(ctx.read_preference(), ReadPreference::Leader)); + assert_eq!(ctx.extension(READ_PREFERENCE_HINT), None); + } else { + assert!(result.is_err()); + } + } + } + #[test] fn test_create_query_context() { let header = RequestHeader {