From a0bb1f75972562424f3d478d8bd5b98b6df970df Mon Sep 17 00:00:00 2001 From: Yang Cen <159225399+BubbleCal@users.noreply.github.com> Date: Thu, 27 Aug 2026 23:03:27 +0800 Subject: [PATCH] fix(functions): harden Rust secret submissions --- .../tests/test_first_class_function_slice2.py | 2 + rust/lancedb/src/function.rs | 153 +++++++++++++++++- rust/lancedb/src/remote/client.rs | 150 +++++++---------- rust/lancedb/src/remote/db.rs | 35 +++- 4 files changed, 243 insertions(+), 97 deletions(-) diff --git a/python/python/tests/test_first_class_function_slice2.py b/python/python/tests/test_first_class_function_slice2.py index 747401f38..983cbf780 100644 --- a/python/python/tests/test_first_class_function_slice2.py +++ b/python/python/tests/test_first_class_function_slice2.py @@ -753,6 +753,7 @@ def test_secret_values_are_validated_before_remote_request( "x" * _MAX_FUNCTION_SECRET_VALUE_BYTES, "é" * (_MAX_FUNCTION_SECRET_VALUE_BYTES // len("é".encode("utf-8"))), ], + ids=["ascii", "multibyte"], ) def test_secret_value_accepts_exact_utf8_byte_limit(value): submission = json.loads(normalize_score._submission_json({"API_TOKEN": value})) @@ -766,6 +767,7 @@ def test_secret_value_accepts_exact_utf8_byte_limit(value): "x" * (_MAX_FUNCTION_SECRET_VALUE_BYTES + 1), "é" * (_MAX_FUNCTION_SECRET_VALUE_BYTES // len("é".encode("utf-8")) + 1), ], + ids=["ascii", "multibyte"], ) def test_secret_value_rejects_over_utf8_byte_limit_before_json_construction( monkeypatch, value diff --git a/rust/lancedb/src/function.rs b/rust/lancedb/src/function.rs index 676533312..2e0be3a88 100644 --- a/rust/lancedb/src/function.rs +++ b/rust/lancedb/src/function.rs @@ -7,7 +7,7 @@ //! This module contains client/wire values only. Catalog persistence, //! environment bake, secret resolution, and execution are owned by Sophon. -use std::collections::BTreeMap; +use std::collections::{BTreeMap, BTreeSet}; use serde::de::{self, DeserializeOwned}; use serde::{Deserialize, Deserializer, Serialize, Serializer}; @@ -15,6 +15,10 @@ use serde_json::Value; use crate::{Error, Result}; +// Keep these byte limits aligned with Sophon's Function submission validation. +pub(crate) const MAX_FUNCTION_SECRET_VALUE_BYTES: usize = 64 * 1024; +const MAX_FUNCTION_SECRET_VALUES_BYTES: usize = 512 * 1024; + fn invalid_json(error: impl std::fmt::Display) -> Error { Error::InvalidInput { message: format!("invalid remote Function JSON: {error}"), @@ -427,6 +431,56 @@ pub struct FunctionRegistrationRequest { pub secret_values: BTreeMap, } +impl FunctionRegistrationRequest { + pub(crate) fn validate_secret_values(&self) -> Result<()> { + let required = self.required_secrets.iter().collect::>(); + let provided = self.secret_values.keys().collect::>(); + if required != provided { + return Err(Error::InvalidInput { + message: "Function secret_values keys must exactly match required_secrets" + .to_string(), + }); + } + + let mut total_bytes = 0usize; + for (name, value) in &self.secret_values { + if value.is_empty() { + return Err(Error::InvalidInput { + message: format!("Function secret {name:?} value must be non-empty"), + }); + } + if value.contains('\0') { + return Err(Error::InvalidInput { + message: format!("Function secret {name:?} value must not contain NUL"), + }); + } + if value.len() > MAX_FUNCTION_SECRET_VALUE_BYTES { + return Err(Error::InvalidInput { + message: format!( + "Function secret {name:?} value exceeds the \ + {MAX_FUNCTION_SECRET_VALUE_BYTES}-byte limit" + ), + }); + } + total_bytes = + total_bytes + .checked_add(value.len()) + .ok_or_else(|| Error::InvalidInput { + message: "Function secret values exceed the request byte limit".to_string(), + })?; + } + if total_bytes > MAX_FUNCTION_SECRET_VALUES_BYTES { + return Err(Error::InvalidInput { + message: format!( + "Function secret values exceed the \ + {MAX_FUNCTION_SECRET_VALUES_BYTES}-byte request limit" + ), + }); + } + Ok(()) + } +} + impl std::fmt::Debug for FunctionRegistrationRequest { fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { let secret_values = self @@ -628,6 +682,103 @@ impl RefreshColumnResult { impl_json!(RefreshColumnResult); +#[cfg(test)] +mod secret_value_tests { + use super::{ + FunctionRegistrationRequest, MAX_FUNCTION_SECRET_VALUE_BYTES, + MAX_FUNCTION_SECRET_VALUES_BYTES, + }; + use crate::Error; + + fn request() -> FunctionRegistrationRequest { + FunctionRegistrationRequest::from_json(include_str!( + "../tests/fixtures/first_class_functions/v1/remote_function_registration_request.json" + )) + .unwrap() + } + + #[test] + fn validates_secret_name_and_value_invariants() { + let missing = request(); + assert!(matches!( + missing.validate_secret_values(), + Err(Error::InvalidInput { message }) if message.contains("exactly match") + )); + + let mut empty = request(); + empty + .secret_values + .insert("API_TOKEN".to_string(), String::new()); + assert!(matches!( + empty.validate_secret_values(), + Err(Error::InvalidInput { message }) if message.contains("non-empty") + )); + + let mut nul = request(); + nul.secret_values + .insert("API_TOKEN".to_string(), "before\0after".to_string()); + assert!(matches!( + nul.validate_secret_values(), + Err(Error::InvalidInput { message }) if message.contains("NUL") + )); + + let mut unexpected = request(); + unexpected + .secret_values + .insert("OTHER".to_string(), "value".to_string()); + assert!(matches!( + unexpected.validate_secret_values(), + Err(Error::InvalidInput { message }) if message.contains("exactly match") + )); + } + + #[test] + fn accepts_exact_secret_value_utf8_byte_limit() { + for value in [ + "x".repeat(MAX_FUNCTION_SECRET_VALUE_BYTES), + "é".repeat(MAX_FUNCTION_SECRET_VALUE_BYTES / "é".len()), + ] { + assert_eq!(value.len(), MAX_FUNCTION_SECRET_VALUE_BYTES); + let mut request = request(); + request.secret_values.insert("API_TOKEN".to_string(), value); + request.validate_secret_values().unwrap(); + } + } + + #[test] + fn rejects_secret_value_over_utf8_byte_limit() { + for value in [ + "x".repeat(MAX_FUNCTION_SECRET_VALUE_BYTES + 1), + "é".repeat(MAX_FUNCTION_SECRET_VALUE_BYTES / "é".len() + 1), + ] { + assert!(value.len() > MAX_FUNCTION_SECRET_VALUE_BYTES); + let mut request = request(); + request.secret_values.insert("API_TOKEN".to_string(), value); + assert!(matches!( + request.validate_secret_values(), + Err(Error::InvalidInput { message }) if message.contains("65536-byte limit") + )); + } + } + + #[test] + fn rejects_aggregate_secret_value_bytes_over_server_limit() { + let mut request = request(); + request.required_secrets = (0..9).map(|index| format!("SECRET_{index}")).collect(); + request.secret_values = request + .required_secrets + .iter() + .map(|name| (name.clone(), "x".repeat(MAX_FUNCTION_SECRET_VALUE_BYTES))) + .collect(); + + assert!(matches!( + request.validate_secret_values(), + Err(Error::InvalidInput { message }) + if message.contains(&format!("{MAX_FUNCTION_SECRET_VALUES_BYTES}-byte request limit")) + )); + } +} + #[cfg(test)] mod conda_environment_tests { use super::PythonEnvironmentSpec; diff --git a/rust/lancedb/src/remote/client.rs b/rust/lancedb/src/remote/client.rs index b4e60dcb2..2e0b410f6 100644 --- a/rust/lancedb/src/remote/client.rs +++ b/rust/lancedb/src/remote/client.rs @@ -45,6 +45,31 @@ fn redacted_json_body(request: &Request) -> Option { serde_json::to_string(&value).ok() } +fn request_log_message(request: &Request, request_id: &str) -> String { + let prefix = format!( + "Sending request_id={}: {} {}", + request_id, + request.method(), + request.url() + ); + let content_type = request + .headers() + .get("content-type") + .and_then(|value| value.to_str().ok()) + .and_then(|value| value.split(';').next()); + if content_type.is_some_and(|value| value.eq_ignore_ascii_case("application/json")) { + // Never format the raw Request here: its Debug representation is not a + // redaction boundary and may include the original body. If the JSON body + // cannot be structurally parsed, suppress it instead of logging raw bytes. + let body = redacted_json_body(request).unwrap_or_else(|| SUPPRESSED_JSON_BODY.to_string()); + format!("{prefix} with body {body}") + } else { + // Method and URL are sufficient request context. Raw Request formatting + // may expose headers or a non-JSON body, so it is never a logging fallback. + prefix + } +} + /// Configuration for TLS/mTLS settings. #[derive(Clone, Debug)] pub struct TlsConfig { @@ -871,27 +896,7 @@ impl RestfulLanceDbClient { pub(crate) fn log_request(&self, request: &Request, request_id: &String) { if log::log_enabled!(log::Level::Debug) { - let content_type = request - .headers() - .get("content-type") - .and_then(|value| value.to_str().ok()) - .and_then(|value| value.split(';').next()); - if content_type.is_some_and(|value| value.eq_ignore_ascii_case("application/json")) { - // Never format the raw Request here: its Debug representation is not a - // redaction boundary and may include the original body. If the JSON body - // cannot be structurally parsed, suppress it instead of logging raw bytes. - let body = - redacted_json_body(request).unwrap_or_else(|| SUPPRESSED_JSON_BODY.to_string()); - debug!( - "Sending request_id={}: {} {} with body {}", - request_id, - request.method(), - request.url(), - body - ); - } else { - debug!("Sending request_id={}: {:?}", request_id, request); - } + debug!("{}", request_log_message(request, request_id)); } } @@ -1102,102 +1107,59 @@ pub mod test_utils { #[cfg(test)] mod tests { use super::*; - use log::{Log, Metadata, Record}; use serial_test::serial; - use std::{ - sync::{ - Once, - atomic::{AtomicBool, Ordering}, - }, - time::Duration, - }; + use std::time::Duration; // Serializes the env-var-mutating tests below: cargo test runs tests in // parallel, but several of these tests read and write the same process- // global env vars (`LANCEDB_USER_ID*`), so they would race without this. static ENV_MUTEX: std::sync::Mutex<()> = std::sync::Mutex::new(()); - static CAPTURED_LOGS: std::sync::Mutex> = std::sync::Mutex::new(Vec::new()); - static CAPTURE_LOGS: AtomicBool = AtomicBool::new(false); - static LOGGER_INIT: Once = Once::new(); - - struct CaptureLogger; - - impl Log for CaptureLogger { - fn enabled(&self, metadata: &Metadata<'_>) -> bool { - metadata.level() <= log::Level::Debug - } - - fn log(&self, record: &Record<'_>) { - if self.enabled(record.metadata()) && CAPTURE_LOGS.load(Ordering::Relaxed) { - CAPTURED_LOGS - .lock() - .unwrap_or_else(|error| error.into_inner()) - .push(format!("{} {}", record.target(), record.args())); - } - } - - fn flush(&self) {} - } - - static CAPTURE_LOGGER: CaptureLogger = CaptureLogger; - - fn initialize_capture_logger() { - LOGGER_INIT.call_once(|| { - log::set_logger(&CAPTURE_LOGGER).expect("test logger should initialize once"); - log::set_max_level(log::LevelFilter::Debug); - }); - } - - fn take_captured_logs() -> String { - let mut logs = CAPTURED_LOGS - .lock() - .unwrap_or_else(|error| error.into_inner()); - let captured = logs.join("\n"); - logs.clear(); - captured - } fn lock_env() -> std::sync::MutexGuard<'static, ()> { ENV_MUTEX.lock().unwrap_or_else(|e| e.into_inner()) } - #[tokio::test] - #[serial(request_logging)] - async fn test_json_request_logging_redacts_secrets_and_suppresses_malformed_bodies() { + #[test] + fn test_request_log_message_redacts_secrets_and_never_formats_raw_requests() { const SECRET_SENTINEL: &str = "udf-secret-log-sentinel-7e4e"; const MALFORMED_SENTINEL: &str = "malformed-secret-log-sentinel-b652"; + const NON_JSON_SENTINEL: &str = "non-json-secret-log-sentinel-7fd1"; - initialize_capture_logger(); - take_captured_logs(); - CAPTURE_LOGS.store(true, Ordering::Relaxed); - - let client = test_utils::client_with_handler(|_| { - http::Response::builder().status(200).body("").unwrap() - }); - let request = client - .post("/v1/functions/create") + let request = reqwest::Client::new() + .post("https://example.com/v1/functions/create") .json(&serde_json::json!({ "name": "uses_secret", "nested": { "secret_values": {"OPENAI_API_KEY": SECRET_SENTINEL}, "safe": "visible-value" } - })); - client.send(request).await.unwrap(); + })) + .build() + .unwrap(); + let log_message = request_log_message(&request, "valid-json"); - let malformed_request = client - .post("/v1/functions/create") + let malformed_request = reqwest::Client::new() + .post("https://example.com/v1/functions/create") .header("content-type", "application/json; charset=utf-8") - .body(format!(r#"{{"secret_values":"{MALFORMED_SENTINEL}""#)); - client.send(malformed_request).await.unwrap(); + .body(format!(r#"{{"secret_values":"{MALFORMED_SENTINEL}""#)) + .build() + .unwrap(); + let malformed_log_message = request_log_message(&malformed_request, "malformed-json"); - CAPTURE_LOGS.store(false, Ordering::Relaxed); - let captured = take_captured_logs(); - assert!(captured.contains("visible-value")); - assert!(captured.contains(REDACTED_JSON_VALUE)); - assert!(captured.contains(SUPPRESSED_JSON_BODY)); - assert!(!captured.contains(SECRET_SENTINEL)); - assert!(!captured.contains(MALFORMED_SENTINEL)); + let non_json_request = reqwest::Client::new() + .post("https://example.com/v1/functions/create") + .header("content-type", "text/plain") + .body(NON_JSON_SENTINEL) + .build() + .unwrap(); + let non_json_log_message = request_log_message(&non_json_request, "non-json"); + + assert!(log_message.contains("visible-value")); + assert!(log_message.contains(REDACTED_JSON_VALUE)); + assert!(!log_message.contains(SECRET_SENTINEL)); + assert!(malformed_log_message.contains(SUPPRESSED_JSON_BODY)); + assert!(!malformed_log_message.contains(MALFORMED_SENTINEL)); + assert!(!non_json_log_message.contains(NON_JSON_SENTINEL)); } #[test] diff --git a/rust/lancedb/src/remote/db.rs b/rust/lancedb/src/remote/db.rs index da9a4b09b..2bca85c2c 100644 --- a/rust/lancedb/src/remote/db.rs +++ b/rust/lancedb/src/remote/db.rs @@ -554,6 +554,7 @@ impl Database for RemoteDatabase { &self, request: FunctionRegistrationRequest, ) -> Result> { + request.validate_secret_values()?; let req = self.client.post("/v1/functions/create").json(&request); let (request_id, response) = self.client.send(req).await?; let response = self.client.check_response(&request_id, response).await?; @@ -2642,7 +2643,8 @@ mod tests { ); const FUNCTION_JOB: &str = include_str!("../../tests/fixtures/first_class_functions/v1/remote_function_job.json"); - let expected: serde_json::Value = serde_json::from_str(REQUEST).unwrap(); + let mut expected: serde_json::Value = serde_json::from_str(REQUEST).unwrap(); + expected["secret_values"] = serde_json::json!({"API_TOKEN": "secret-value"}); let conn = Connection::new_with_handler(move |request| match request.url().path() { "/v1/functions/create" => { assert_eq!(request.method(), &reqwest::Method::POST); @@ -2660,7 +2662,10 @@ mod tests { .unwrap(), path => panic!("unexpected path: {path}"), }); - let request = crate::function::FunctionRegistrationRequest::from_json(REQUEST).unwrap(); + let mut request = crate::function::FunctionRegistrationRequest::from_json(REQUEST).unwrap(); + request + .secret_values + .insert("API_TOKEN".to_string(), "secret-value".to_string()); let job = conn.create_function_async(request).await.unwrap(); assert_eq!(job.id(), Some("job-function-1")); let version = job.wait().await.unwrap(); @@ -2668,6 +2673,32 @@ mod tests { assert_eq!(version.version(), "fv_01K3EXACT"); } + #[tokio::test] + async fn test_create_function_async_validates_secrets_before_serialization_and_send() { + const REQUEST: &str = include_str!( + "../../tests/fixtures/first_class_functions/v1/remote_function_registration_request.json" + ); + let sends = Arc::new(AtomicUsize::new(0)); + let sends_ref = sends.clone(); + let conn = Connection::new_with_handler(move |_| { + sends_ref.fetch_add(1, Ordering::SeqCst); + http::Response::builder().status(500).body("").unwrap() + }); + let mut request = crate::function::FunctionRegistrationRequest::from_json(REQUEST).unwrap(); + request.secret_values.insert( + "API_TOKEN".to_string(), + "x".repeat(crate::function::MAX_FUNCTION_SECRET_VALUE_BYTES + 1), + ); + + let error = conn.create_function_async(request).await.unwrap_err(); + + assert!(matches!( + error, + Error::InvalidInput { message } if message.contains("65536-byte limit") + )); + assert_eq!(sends.load(Ordering::SeqCst), 0); + } + #[tokio::test] async fn test_get_function_requires_and_sends_exact_version() { const VERSION: &str = include_str!(