diff --git a/python/python/lancedb/functions.py b/python/python/lancedb/functions.py index b2a794b00..0e0e6bccf 100644 --- a/python/python/lancedb/functions.py +++ b/python/python/lancedb/functions.py @@ -486,6 +486,8 @@ class RefreshColumnResult(_RemoteValue): _FUNCTION_NAME = re.compile(r"^[A-Za-z_][A-Za-z0-9_.-]*$") _SECRET_NAME = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$") +# Keep this byte limit aligned with Sophon's MAX_FUNCTION_SECRET_VALUE_BYTES. +_MAX_FUNCTION_SECRET_VALUE_BYTES = 64 * 1024 def _validate_secret_value(name: str, value: Any) -> str: @@ -496,6 +498,12 @@ def _validate_secret_value(name: str, value: Any) -> str: raise ValueError(f"Function secret {name!r} value must be non-empty") if "\0" in value: raise ValueError(f"Function secret {name!r} value must not contain NUL") + value_bytes = len(value.encode("utf-8")) + if value_bytes > _MAX_FUNCTION_SECRET_VALUE_BYTES: + raise ValueError( + f"Function secret {name!r} value exceeds the " + f"{_MAX_FUNCTION_SECRET_VALUE_BYTES}-byte limit" + ) return value diff --git a/python/python/tests/test_first_class_function_slice2.py b/python/python/tests/test_first_class_function_slice2.py index b3f67cb49..747401f38 100644 --- a/python/python/tests/test_first_class_function_slice2.py +++ b/python/python/tests/test_first_class_function_slice2.py @@ -19,7 +19,12 @@ import pyarrow as pa import pytest import lancedb -from lancedb.functions import FunctionRegistrationRequest, UdfDefinition, udf +from lancedb.functions import ( + _MAX_FUNCTION_SECRET_VALUE_BYTES, + FunctionRegistrationRequest, + UdfDefinition, + udf, +) THRESHOLD = 20 _CACHE = None @@ -742,6 +747,39 @@ def test_secret_values_are_validated_before_remote_request( assert state["requests"] == [] +@pytest.mark.parametrize( + "value", + [ + "x" * _MAX_FUNCTION_SECRET_VALUE_BYTES, + "é" * (_MAX_FUNCTION_SECRET_VALUE_BYTES // len("é".encode("utf-8"))), + ], +) +def test_secret_value_accepts_exact_utf8_byte_limit(value): + submission = json.loads(normalize_score._submission_json({"API_TOKEN": value})) + assert submission["secret_values"]["API_TOKEN"] == value + assert len(value.encode("utf-8")) == _MAX_FUNCTION_SECRET_VALUE_BYTES + + +@pytest.mark.parametrize( + "value", + [ + "x" * (_MAX_FUNCTION_SECRET_VALUE_BYTES + 1), + "é" * (_MAX_FUNCTION_SECRET_VALUE_BYTES // len("é".encode("utf-8")) + 1), + ], +) +def test_secret_value_rejects_over_utf8_byte_limit_before_json_construction( + monkeypatch, value +): + def fail_if_json_construction_starts(self): + pytest.fail("oversized secret reached JSON construction") + + monkeypatch.setattr( + FunctionRegistrationRequest, "_known_dict", fail_if_json_construction_starts + ) + with pytest.raises(ValueError, match=r"exceeds the 65536-byte limit"): + normalize_score._submission_json({"API_TOKEN": value}) + + @pytest.mark.asyncio async def test_async_remote_registration_submits_secret_values_only_once(): with _mock_remote_function_catalog() as (host, state): diff --git a/rust/lancedb/src/remote/client.rs b/rust/lancedb/src/remote/client.rs index 57dd89890..b4e60dcb2 100644 --- a/rust/lancedb/src/remote/client.rs +++ b/rust/lancedb/src/remote/client.rs @@ -7,6 +7,7 @@ use reqwest::{ Body, Request, RequestBuilder, Response, header::{HeaderMap, HeaderValue}, }; +use serde_json::Value; use std::{collections::HashMap, future::Future, str::FromStr, sync::Arc, time::Duration}; use crate::error::{Error, Result}; @@ -14,6 +15,35 @@ use crate::remote::db::RemoteOptions; use crate::remote::retry::{ResolvedRetryConfig, RetryCounter}; const REQUEST_ID_HEADER: HeaderName = HeaderName::from_static("x-request-id"); +const REDACTED_JSON_VALUE: &str = "[REDACTED]"; +const SUPPRESSED_JSON_BODY: &str = "[JSON BODY SUPPRESSED]"; + +fn is_sensitive_json_field(name: &str) -> bool { + name.to_ascii_lowercase().contains("secret") +} + +fn redact_sensitive_json_fields(value: &mut Value) { + match value { + Value::Object(fields) => { + for (name, child) in fields { + if is_sensitive_json_field(name) { + *child = Value::String(REDACTED_JSON_VALUE.to_string()); + } else { + redact_sensitive_json_fields(child); + } + } + } + Value::Array(values) => values.iter_mut().for_each(redact_sensitive_json_fields), + _ => {} + } +} + +fn redacted_json_body(request: &Request) -> Option { + let body = request.body()?.as_bytes()?; + let mut value = serde_json::from_slice(body).ok()?; + redact_sensitive_json_fields(&mut value); + serde_json::to_string(&value).ok() +} /// Configuration for TLS/mTLS settings. #[derive(Clone, Debug)] @@ -844,13 +874,20 @@ impl RestfulLanceDbClient { let content_type = request .headers() .get("content-type") - .map(|v| v.to_str().unwrap()); - if content_type == Some("application/json") { - let body = request.body().as_ref().unwrap().as_bytes().unwrap(); - let body = String::from_utf8_lossy(body); + .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, body + "Sending request_id={}: {} {} with body {}", + request_id, + request.method(), + request.url(), + body ); } else { debug!("Sending request_id={}: {:?}", request_id, request); @@ -1065,18 +1102,104 @@ pub mod test_utils { #[cfg(test)] mod tests { use super::*; + use log::{Log, Metadata, Record}; use serial_test::serial; - use std::time::Duration; + use std::{ + sync::{ + Once, + atomic::{AtomicBool, Ordering}, + }, + 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() { + const SECRET_SENTINEL: &str = "udf-secret-log-sentinel-7e4e"; + const MALFORMED_SENTINEL: &str = "malformed-secret-log-sentinel-b652"; + + 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") + .json(&serde_json::json!({ + "name": "uses_secret", + "nested": { + "secret_values": {"OPENAI_API_KEY": SECRET_SENTINEL}, + "safe": "visible-value" + } + })); + client.send(request).await.unwrap(); + + let malformed_request = client + .post("/v1/functions/create") + .header("content-type", "application/json; charset=utf-8") + .body(format!(r#"{{"secret_values":"{MALFORMED_SENTINEL}""#)); + client.send(malformed_request).await.unwrap(); + + 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)); + } + #[test] fn test_timeout_config_default() { let config = TimeoutConfig::default();