From 68a12b79af290b7e8df78f811c9c790d89d226d0 Mon Sep 17 00:00:00 2001 From: Dmitri Plotnikov Date: Sat, 26 Sep 2026 21:22:13 -0700 Subject: [PATCH] [Python] Fix AB-BA deadlock between Python GIL and DescriptorPool mutex Ensure all C++ locks (Program mutex, Activation mutex, once flags, DescriptorPool mutex) precede Python GIL acquisition by releasing the GIL in PyCelExpression::Eval, PyObjectToCelValue, and PyCelEnvInternal constructor, while acquiring the GIL on callbacks into Python. PiperOrigin-RevId: 989033642 --- cel_expr_python/cel_parallel_test.py | 49 ++++++++-- cel_expr_python/cel_test.py | 81 +++++++++++++++- cel_expr_python/py_cel_env_internal.cc | 41 ++++---- cel_expr_python/py_cel_env_internal.h | 2 +- cel_expr_python/py_cel_expression.cc | 37 ++++--- cel_expr_python/py_cel_expression.h | 10 +- cel_expr_python/py_cel_function.cc | 31 ++++-- cel_expr_python/py_cel_function.h | 10 +- cel_expr_python/py_cel_overload.h | 2 +- cel_expr_python/py_cel_value.cc | 112 ++++++++++++++-------- cel_expr_python/py_descriptor_database.cc | 3 + 11 files changed, 281 insertions(+), 97 deletions(-) diff --git a/cel_expr_python/cel_parallel_test.py b/cel_expr_python/cel_parallel_test.py index 5a41fbe..240db88 100644 --- a/cel_expr_python/cel_parallel_test.py +++ b/cel_expr_python/cel_parallel_test.py @@ -106,6 +106,18 @@ def setUp(self): "var_int_map": cel.Type.Map(cel.Type.INT, cel.Type.STRING), "var_msg": cel.Type("cel.expr.conformance.proto2.TestAllTypes"), }, + functions=[ + cel.FunctionDecl( + "custom_fn", + [ + cel.Overload( + "custom_fn_int", + return_type=cel.Type.INT, + parameters=[cel.Type.INT], + ) + ], + ) + ], ) self.object_counts_before_test = self._grab_object_counts() @@ -170,28 +182,28 @@ def testSequentialEval(self): self._test_eval(multi_threaded=False) def _test_compile(self, multi_threaded: bool): - def compile_expr(n: int) -> cel.Expression: + + def compile_and_eval(n: int) -> Any: test_case = _TEST_CASES[n % len(_TEST_CASES)] - return self.env.compile(test_case.expr) + expr = self.env.compile(test_case.expr) + data = test_case.data(n) + return expr.eval(data=data).plain_value() start_time = time.perf_counter() if multi_threaded: with concurrent.futures.ThreadPoolExecutor(max_workers=8) as executor: - results = list(executor.map(compile_expr, range(_NUM_COMPILATIONS))) + results = list(executor.map(compile_and_eval, range(_NUM_COMPILATIONS))) else: - results = [compile_expr(n) for n in range(_NUM_COMPILATIONS)] + results = [compile_and_eval(n) for n in range(_NUM_COMPILATIONS)] duration_ms = (time.perf_counter() - start_time) * 1000 mode = "Multi-threaded" if multi_threaded else "Sequential" logging.info("%s compilation duration: %.2f ms", mode, duration_ms) self.assertLen(results, _NUM_COMPILATIONS) - for i, expr in enumerate(results): + for i, res in enumerate(results): test_case = _TEST_CASES[i % len(_TEST_CASES)] - data = test_case.data(i) - self.assertEqual( - expr.eval(data=data).plain_value(), test_case.expected(i) - ) + self.assertEqual(res, test_case.expected(i)) def testMultiThreadedCompilation(self): self._test_compile(multi_threaded=True) @@ -240,6 +252,25 @@ def read_values(_): for i in range(10): run_concurrent_value_test(i) + def testConcurrentCustomFunction(self): + expr = self.env.compile("custom_fn(var_int)") + + def eval_custom_fn(n: int): + fn = cel.Function( + "custom_fn", + [cel.Type.INT], + False, + lambda x: x * 2, + return_type=cel.Type.INT, + ) + act = self.env.Activation({"var_int": n}, functions=[fn]) + res = expr.eval(act) + self.assertEqual(res.value(), n * 2) + + with concurrent.futures.ThreadPoolExecutor(max_workers=8) as executor: + results = list(executor.map(eval_custom_fn, range(100))) + self.assertLen(results, 100) + if __name__ == "__main__": absltest.main() diff --git a/cel_expr_python/cel_test.py b/cel_expr_python/cel_test.py index 861a2d2..fd5c8ef 100644 --- a/cel_expr_python/cel_test.py +++ b/cel_expr_python/cel_test.py @@ -846,6 +846,45 @@ def testErrorHandling(self): r"Could not find file containing symbol:.* \[NOT_FOUND\]", ) + def testDescriptorPoolMissingSerializedPb(self): + bad_env = cel.NewEnv( + _MissingSerializedPbPool(), + variables={}, + options=self.options, + ) + with self.assertRaises(Exception) as e: + bad_env.compile("cel.expr.conformance.proto2.TestSomeTypes{}") + self.assertIn( + "Python file has no attribute 'serialized_pb'", + str(e.exception), + ) + + def testDescriptorPoolCorruptedSerializedPb(self): + bad_env = cel.NewEnv( + _CorruptedSerializedPbPool(), + variables={}, + options=self.options, + ) + with self.assertRaises(Exception) as e: + bad_env.compile("cel.expr.conformance.proto2.TestSomeTypes{}") + self.assertIn( + "Failed to parse descriptor for" + " cel.expr.conformance.proto2.TestSomeTypes", + str(e.exception), + ) + + def testProtoMessageToCelValueError(self): + bad_env = cel.NewEnv( + _RaisingDescriptorPool(), + variables={"var_proto": cel.Type.DYN}, + options=self.options, + ) + expr = bad_env.compile("var_proto", disable_check=True) + msg = test_all_types_pb.TestAllTypes(single_string="Hey") + res = expr.eval(data={"var_proto": msg}) + self.assertEqual(res.type(), cel.Type.ERROR) + self.assertIn("Custom pool error", str(res.value())) + class CompatibleNumber: @@ -872,7 +911,47 @@ def __int__(self) -> int: class _BadDescriptorPool: def FindFileContainingSymbol(self, symbol_name: str): # pylint: disable=invalid-name - raise LookupError("Could not find file containing symbol: %s" % symbol_name) + if symbol_name.startswith("cel.expr.conformance"): + raise LookupError( + "Could not find file containing symbol: %s" % symbol_name + ) + raise KeyError(symbol_name) + + +class _MissingSerializedPbPool: + + def FindFileByName( # pylint: disable=invalid-name,unused-argument + self, filename: str + ): + raise KeyError(filename) + + def FindFileContainingSymbol( # pylint: disable=invalid-name,unused-argument + self, symbol_name: str + ): + if symbol_name.startswith("cel.expr.conformance"): + return object() + raise KeyError(symbol_name) + + +class _CorruptedSerializedPbPool: + + class _FakeFile: + serialized_pb = b"corrupted proto descriptor bytes" + + def FindFileContainingSymbol( # pylint: disable=invalid-name,unused-argument + self, symbol_name: str + ): + if symbol_name.startswith("cel.expr.conformance"): + return self._FakeFile() + raise KeyError(symbol_name) + + +class _RaisingDescriptorPool: + + def FindFileContainingSymbol(self, symbol_name: str): # pylint: disable=invalid-name + if symbol_name.startswith("cel.expr.conformance"): + raise RuntimeError("Custom pool error: %s" % symbol_name) + raise KeyError(symbol_name) class CelWithoutProtoSupportTest(absltest.TestCase): diff --git a/cel_expr_python/py_cel_env_internal.cc b/cel_expr_python/py_cel_env_internal.cc index 4bb32f8..7983bfe 100644 --- a/cel_expr_python/py_cel_env_internal.cc +++ b/cel_expr_python/py_cel_env_internal.cc @@ -115,11 +115,16 @@ PyCelEnvInternal::PyCelEnvInternal( google::protobuf::Arena arena; for (const cel::Config::VariableConfig& variable_config : env_config_.GetConfig().GetVariableConfigs()) { - auto status_or_type = cel::TypeInfoToType(variable_config.type_info, - descriptor_pool_.get(), &arena); - if (status_or_type.ok()) { - variable_types_[variable_config.name] = - PyCelType::FromCelType(*status_or_type); + absl::StatusOr type; + { + // Release the GIL during TypeInfoToType lookups to prevent lock order + // inversion with DescriptorPool's internal mutex. + py::gil_scoped_release gil_release; + type = cel::TypeInfoToType(variable_config.type_info, + descriptor_pool_.get(), &arena); + } + if (type.ok()) { + variable_types_[variable_config.name] = PyCelType::FromCelType(*type); } } } @@ -132,7 +137,7 @@ PyCelEnvInternal::NewCelEnvInternal( PyObject* py_descriptor_pool, const std::unordered_map& variable_types, const std::vector& extensions, - cel::ExpressionContainer container, + const cel::ExpressionContainer& container, const std::vector>& functions, const std::unordered_map& function_impls) { cel::Config config = env_config.GetConfig(); @@ -310,17 +315,6 @@ absl::StatusOr> PyCelEnvInternal::BuildRuntime( GetEnvConfig().GetConfig().GetFunctionConfigs()) { for (const cel::Config::FunctionOverloadConfig& overload_config : function_config.overload_configs) { - auto it = function_impls_.find(overload_config.overload_id); - if (it == function_impls_.end()) { - continue; - } - py::object py_function; - if (!PyGILState_Check()) { - py::gil_scoped_acquire acquire; - py_function = it->second; - } else { - py_function = it->second; - } std::vector param_kinds; param_kinds.reserve(overload_config.parameters.size()); for (const cel::Config::TypeInfo& parameter : @@ -333,15 +327,20 @@ absl::StatusOr> PyCelEnvInternal::BuildRuntime( cel::FunctionDescriptor descriptor( function_config.name, overload_config.is_member_function, param_kinds, kFunctionDescriptorOptions); + auto it = function_impls_.find(overload_config.overload_id); + if (it == function_impls_.end()) { + CEL_PYTHON_RETURN_IF_ERROR( + builder.function_registry().RegisterLazyFunction(descriptor)); + continue; + } CEL_PYTHON_ASSIGN_OR_RETURN( cel::Type return_type, cel::TypeInfoToType(overload_config.return_type, descriptor_pool_.get(), &arena)); CEL_PYTHON_RETURN_IF_ERROR(builder.function_registry().Register( - descriptor, - std::make_unique( - function_config.name, PyCelType::FromCelType(return_type), - std::move(py_function)))); + descriptor, std::make_unique( + function_config.name, + PyCelType::FromCelType(return_type), it->second))); } } return std::move(builder).Build(); diff --git a/cel_expr_python/py_cel_env_internal.h b/cel_expr_python/py_cel_env_internal.h index 95bfd14..4a0b35c 100644 --- a/cel_expr_python/py_cel_env_internal.h +++ b/cel_expr_python/py_cel_env_internal.h @@ -76,7 +76,7 @@ class PyCelEnvInternal { PyObject* py_descriptor_pool, const std::unordered_map& variable_types, const std::vector& extensions, - cel::ExpressionContainer container, + const cel::ExpressionContainer& container, const std::vector>& functions, const std::unordered_map& function_impls); diff --git a/cel_expr_python/py_cel_expression.cc b/cel_expr_python/py_cel_expression.cc index 53084d4..41a645a 100644 --- a/cel_expr_python/py_cel_expression.cc +++ b/cel_expr_python/py_cel_expression.cc @@ -33,6 +33,7 @@ #include "absl/status/status.h" #include "absl/status/statusor.h" #include "absl/strings/str_format.h" +#include "absl/synchronization/mutex.h" #include "checker/validation_result.h" #include "common/ast.h" #include "common/ast_proto.h" @@ -43,7 +44,6 @@ #include "parser/parser_interface.h" #include "runtime/embedder_context.h" #include "runtime/runtime.h" -#include "cel_expr_python/free_threading_mutex.h" #include "cel_expr_python/py_cel_activation.h" #include "cel_expr_python/py_cel_arena.h" #include "cel_expr_python/py_cel_env_internal.h" @@ -129,7 +129,8 @@ absl::StatusOr PyCelExpression::Compile( } PyCelExpression::PyCelExpression(PyCelExpression&& other) noexcept { - FreeThreadingLockGuard lock(other.mutex_); + // Lock other.mutex_ to safely move the Program and expression state. + absl::MutexLock lock(other.mutex_); expr_ = std::move(other.expr_); env_ = std::move(other.env_); cel_program_ = std::move(other.cel_program_); @@ -155,7 +156,10 @@ PyCelType PyCelExpression::GetReturnType() { } absl::StatusOr PyCelExpression::GetProgram() { - FreeThreadingLockGuard lock(mutex_); + // Lock is needed in both GIL-enabled and free-threaded builds because Eval() + // releases the Python GIL before calling GetProgram(), allowing multiple + // threads to concurrently initialize cel_program_. + absl::MutexLock lock(mutex_); if (cel_program_) { return cel_program_.get(); } @@ -179,17 +183,26 @@ absl::StatusOr PyCelExpression::GetProgram() { absl::StatusOr PyCelExpression::Eval( const PyCelActivation& activation) { ABSL_CHECK(PyGILState_Check()); - CEL_PYTHON_ASSIGN_OR_RETURN(const cel::Program* program, GetProgram()); std::shared_ptr arena = activation.GetArena(); std::shared_ptr env = activation.GetEnv(); - cel::EmbedderContext embedder_context = cel::EmbedderContext::From(&env); - cel::EvaluateOptions options; - options.message_factory = env->GetMessageFactory(); - options.embedder_context = &embedder_context; - CEL_PYTHON_ASSIGN_OR_RETURN( - cel::Value result, - program->Evaluate(arena->GetArena(), *activation.GetActivation(), - std::move(options))); + cel::Value result; + { + // Release the GIL before entering C++ program creation and evaluation to + // prevent lock inversion/deadlock with DescriptorPool's internal mutex + // (C++ Locks -> DescriptorPool Mutex -> Python GIL). Callbacks into Python + // (such as PyCelValueProvider::Provide and PyCelFunctionAdapter::Invoke) + // re-acquire the GIL on demand. + py::gil_scoped_release gil_release; + CEL_PYTHON_ASSIGN_OR_RETURN(const cel::Program* program, GetProgram()); + cel::EmbedderContext embedder_context = cel::EmbedderContext::From(&env); + cel::EvaluateOptions options; + options.message_factory = env->GetMessageFactory(); + options.embedder_context = &embedder_context; + CEL_PYTHON_ASSIGN_OR_RETURN( + result, + program->Evaluate(arena->GetArena(), *activation.GetActivation(), + std::move(options))); + } return PyCelValue(result, arena, std::move(env)); } diff --git a/cel_expr_python/py_cel_expression.h b/cel_expr_python/py_cel_expression.h index acc1287..eaeb369 100644 --- a/cel_expr_python/py_cel_expression.h +++ b/cel_expr_python/py_cel_expression.h @@ -24,8 +24,8 @@ #include "cel/expr/syntax.pb.h" #include "absl/base/thread_annotations.h" #include "absl/status/statusor.h" +#include "absl/synchronization/mutex.h" #include "runtime/runtime.h" -#include "cel_expr_python/free_threading_mutex.h" #include "cel_expr_python/py_cel_activation.h" #include "cel_expr_python/py_cel_type.h" #include "cel_expr_python/py_cel_value.h" @@ -70,7 +70,13 @@ class PyCelExpression { std::variant expr_; std::shared_ptr env_; - mutable FreeThreadingMutex mutex_; + // mutex_ protects lazy initialization of cel_program_ in GetProgram(). + // Unlike FreeThreadingMutex (which compiles to a no-op in GIL builds), a real + // absl::Mutex is required here because Eval() releases the Python GIL before + // calling GetProgram(), allowing multiple threads in both GIL-enabled and + // free-threaded Python builds to concurrently access and initialize + // cel_program_. + mutable absl::Mutex mutex_; std::unique_ptr cel_program_ ABSL_GUARDED_BY(mutex_); }; diff --git a/cel_expr_python/py_cel_function.cc b/cel_expr_python/py_cel_function.cc index 5ddd2b4..70f83ae 100644 --- a/cel_expr_python/py_cel_function.cc +++ b/cel_expr_python/py_cel_function.cc @@ -75,15 +75,33 @@ PyCelFunction::PyCelFunction(std::string function_name, PyCelFunctionAdapter::PyCelFunctionAdapter(std::string function_name, PyCelType return_type, - py::object py_function) + const py::object& py_function) : function_name_(std::move(function_name)), return_type_(std::move(return_type)), - py_function_(std::move(py_function)) {} + py_function_(py_function.ptr()) { + if (!PyGILState_Check()) { + py::gil_scoped_acquire acquire; + Py_XINCREF(py_function_); + } else { + Py_XINCREF(py_function_); + } +} + +PyCelFunctionAdapter::~PyCelFunctionAdapter() { + if (py_function_ != nullptr) { + if (!PyGILState_Check()) { + py::gil_scoped_acquire acquire; + Py_XDECREF(py_function_); + } else { + Py_XDECREF(py_function_); + } + } +} absl::StatusOr PyCelFunctionAdapter::Invoke( absl::Span args, const cel::Function::InvokeContext& context) const { - ABSL_CHECK(PyGILState_Check()); + py::gil_scoped_acquire acquire; std::shared_ptr env = GetEnvFromContext(context); CEL_PYTHON_ASSIGN_OR_RETURN(auto py_arena, @@ -94,7 +112,7 @@ absl::StatusOr PyCelFunctionAdapter::Invoke( CelValueToPyObject(args[i], env, py_arena, /*plain_value=*/true)); } - PyObject* result = PyObject_CallObject(py_function_.ptr(), py_args); + PyObject* result = PyObject_CallObject(py_function_, py_args); Py_DECREF(py_args); absl::Status status = PyErr_toStatus(); if (!status.ok()) { @@ -105,9 +123,8 @@ absl::StatusOr PyCelFunctionAdapter::Invoke( absl::StatusOr cel_result = PyObjectToCelValue( result, return_type_, [this]() { - return absl::StrFormat( - "Python function '%s'", - PyUnicode_AsUTF8(PyObject_Repr(py_function_.ptr()))); + return absl::StrFormat("Python function '%s'", + PyUnicode_AsUTF8(PyObject_Repr(py_function_))); }, env, context.arena()); Py_XDECREF(result); diff --git a/cel_expr_python/py_cel_function.h b/cel_expr_python/py_cel_function.h index 87dfb5c..6a28112 100644 --- a/cel_expr_python/py_cel_function.h +++ b/cel_expr_python/py_cel_function.h @@ -45,7 +45,7 @@ class PyCelFunction { std::string function_name() const { return function_name_; } const std::vector& parameters() const { return parameters_; } bool is_member() const { return is_member_; } - py::object impl() const { return impl_; } + const py::object& impl() const { return impl_; } const PyCelType& return_type() const { return return_type_; } private: @@ -61,7 +61,11 @@ class PyCelFunction { class PyCelFunctionAdapter : public cel::Function { public: PyCelFunctionAdapter(std::string function_name, PyCelType return_type, - py::object py_function); + const py::object& py_function); + ~PyCelFunctionAdapter() override; + + PyCelFunctionAdapter(const PyCelFunctionAdapter&) = delete; + PyCelFunctionAdapter& operator=(const PyCelFunctionAdapter&) = delete; absl::StatusOr Invoke( absl::Span args, @@ -70,7 +74,7 @@ class PyCelFunctionAdapter : public cel::Function { private: std::string function_name_; PyCelType return_type_; - py::object py_function_; + PyObject* py_function_; }; } // namespace cel_python diff --git a/cel_expr_python/py_cel_overload.h b/cel_expr_python/py_cel_overload.h index 05d4e16..3218b08 100644 --- a/cel_expr_python/py_cel_overload.h +++ b/cel_expr_python/py_cel_overload.h @@ -43,7 +43,7 @@ class PyCelOverload { PyCelType return_type() const { return return_type_; } const std::vector& parameters() const { return parameters_; } bool is_member() const { return is_member_; } - py::object py_function() const { return py_function_; } + const py::object& py_function() const { return py_function_; } cel::Config::FunctionOverloadConfig ToFunctionOverloadConfig() const; diff --git a/cel_expr_python/py_cel_value.cc b/cel_expr_python/py_cel_value.cc index df26c46..f0c78d5 100644 --- a/cel_expr_python/py_cel_value.cc +++ b/cel_expr_python/py_cel_value.cc @@ -166,7 +166,7 @@ PyCelValueProvider::~PyCelValueProvider() { cel::Value PyCelValueProvider::Provide( const google::protobuf::DescriptorPool* descriptor_pool, google::protobuf::MessageFactory* message_factory, google::protobuf::Arena* arena) const { - ABSL_CHECK(PyGILState_Check()); + py::gil_scoped_acquire acquire; const PyCelType& type = env_->GetVariableType(name_); absl::StatusOr converted_value = PyObjectToCelValue( py_object_, type, [this]() { return name_; }, env_, arena); @@ -510,6 +510,74 @@ static void EnsureDateTimeModuleImported() { absl::call_once(import_py_datetime_once_flag, []() { PyDateTime_IMPORT; }); } +// Converts a serialized protocol buffer message to a CEL value. +// +// The CEL value of the protocol buffer message or cel::ErrorValue if the +// conversion fails. +static cel::Value ProtoMessageToCelValue( + absl::string_view message_type_name, PyObject* serialized_bytes, + const std::shared_ptr& env, google::protobuf::Arena* arena) { + // The Python GIL is required to safely inspect the serialized_bytes PyObject. + // The caller retains ownership of serialized_bytes, so the extracted buffer + // (bytes_ptr and bytes_size) remains valid across GIL release within the + // lifetime of this function. + const uint8_t* bytes_ptr = + reinterpret_cast(PyBytes_AS_STRING(serialized_bytes)); + const Py_ssize_t bytes_size = PyBytes_GET_SIZE(serialized_bytes); + + absl::StatusOr wrapped_message; + { + // Release the GIL before calling into C++ DescriptorPool and MessageFactory + // to maintain the lock hierarchy (DescriptorPool Mutex -> Python GIL). + // Any fallback lookups into Python (e.g., via PyDescriptorDatabase) + // will re-acquire the GIL on demand. + py::gil_scoped_release gil_release; + + const google::protobuf::Descriptor* descriptor = + env->GetDescriptorPool()->FindMessageTypeByName(message_type_name); + if (descriptor == nullptr) { + wrapped_message = absl::InvalidArgumentError(absl::StrFormat( + "Descriptor not found for message type '%s'", message_type_name)); + } else { + const google::protobuf::Message* prototype = + env->GetMessageFactory()->GetPrototype(descriptor); + if (prototype == nullptr) { + wrapped_message = absl::InvalidArgumentError(absl::StrFormat( + "Prototype not found for message type '%s'", message_type_name)); + } else { + google::protobuf::Message* message = prototype->New(arena); + if (message == nullptr) { + wrapped_message = absl::InternalError(absl::StrFormat( + "Failed to create new message of type '%s'", message_type_name)); + } else { + google::protobuf::io::CodedInputStream coded_input_stream(bytes_ptr, + bytes_size); + if (!message->MergePartialFromCodedStream(&coded_input_stream)) { + wrapped_message = absl::InvalidArgumentError(absl::StrFormat( + "Failed to parse serialized data for type '%s' ", + message_type_name)); + } else { + wrapped_message = + cel::Value::WrapMessage(message, env->GetDescriptorPool(), + env->GetMessageFactory(), arena); + } + } + } + } + } + + // The GIL is re-acquired here. Check if protobuf parsing or descriptor + // lookups raised a Python exception in PyDescriptorDatabase. + absl::Status status = PyErr_toStatus(); + if (!status.ok()) { + return cel::ErrorValue(status); + } + if (!wrapped_message.ok()) { + return cel::ErrorValue(wrapped_message.status()); + } + return *wrapped_message; +} + absl::StatusOr PyObjectToCelValue( PyObject* py_object, const PyCelType& expected_type, absl::FunctionRef context, @@ -741,46 +809,10 @@ absl::StatusOr PyObjectToCelValue( return InvalidTypeError(py_object, context, expected_type); } - const std::string& message_type_name = type.ToString(); - const google::protobuf::Descriptor* descriptor = - env->GetDescriptorPool()->FindMessageTypeByName(message_type_name); - if (descriptor == nullptr) { - return cel::ErrorValue(absl::InvalidArgumentError(absl::StrFormat( - "Descriptor not found for message type '%s'", message_type_name))); - } - - const google::protobuf::Message* prototype = - env->GetMessageFactory()->GetPrototype(descriptor); - if (prototype == nullptr) { - return cel::ErrorValue(absl::InvalidArgumentError(absl::StrFormat( - "Prototype not found for message type '%s'", message_type_name))); - } - google::protobuf::Message* message = prototype->New(arena); - if (message == nullptr) { - return cel::ErrorValue(absl::InternalError(absl::StrFormat( - "Failed to create new message of type '%s'", message_type_name))); - } - - // Create a CodedInputStream to read the serialized bytes directly from - // the Python bytes object, without copying. - google::protobuf::io::CodedInputStream coded_input_stream( - reinterpret_cast(PyBytes_AS_STRING(serialized_bytes)), - PyBytes_GET_SIZE(serialized_bytes)); - if (!message->MergePartialFromCodedStream(&coded_input_stream)) { - Py_DECREF(serialized_bytes); - return cel::ErrorValue(absl::InvalidArgumentError( - absl::StrFormat("Failed to parse serialized data for type '%s' ", - message_type_name))); - } + cel::Value value = + ProtoMessageToCelValue(type.ToString(), serialized_bytes, env, arena); Py_DECREF(serialized_bytes); - // Protobuf parsing may have run into a Python exception in - // PyDescriptorDatabase, which makes Python native calls. - absl::Status status = PyErr_toStatus(); - if (!status.ok()) { - return cel::ErrorValue(status); - } - return cel::Value::WrapMessage(message, env->GetDescriptorPool(), - env->GetMessageFactory(), arena); + return value; } case cel::Kind::kList: { if (PyList_Check(py_object)) { diff --git a/cel_expr_python/py_descriptor_database.cc b/cel_expr_python/py_descriptor_database.cc index bbbb29f..95148b2 100644 --- a/cel_expr_python/py_descriptor_database.cc +++ b/cel_expr_python/py_descriptor_database.cc @@ -83,6 +83,7 @@ bool PyDescriptorDatabase::FindFileByName(StringViewArg filename, if (pyfile_serialized == nullptr) { PyErr_Format(PyExc_TypeError, "Python file has no attribute 'serialized_pb'"); + PyErr_noteAndClear(); return false; } @@ -133,6 +134,7 @@ bool PyDescriptorDatabase::FindFileContainingSymbol( if (pyfile_serialized == nullptr) { PyErr_Format(PyExc_TypeError, "Python file has no attribute 'serialized_pb'"); + PyErr_noteAndClear(); return false; } @@ -142,6 +144,7 @@ bool PyDescriptorDatabase::FindFileContainingSymbol( if (!ok) { PyErr_Format(PyExc_RuntimeError, "Failed to parse descriptor for %s", symbol_name.data()); + PyErr_noteAndClear(); } Py_DECREF(pyfile_serialized); return ok;