Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
49 changes: 40 additions & 9 deletions cel_expr_python/cel_parallel_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()

Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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()
81 changes: 80 additions & 1 deletion cel_expr_python/cel_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:

Expand All @@ -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):
Expand Down
41 changes: 20 additions & 21 deletions cel_expr_python/py_cel_env_internal.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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<cel::Type> 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);
}
}
}
Expand All @@ -132,7 +137,7 @@ PyCelEnvInternal::NewCelEnvInternal(
PyObject* py_descriptor_pool,
const std::unordered_map<std::string, PyCelType>& variable_types,
const std::vector<PyObject*>& extensions,
cel::ExpressionContainer container,
const cel::ExpressionContainer& container,
const std::vector<std::shared_ptr<PyCelFunctionDecl>>& functions,
const std::unordered_map<std::string, py::object>& function_impls) {
cel::Config config = env_config.GetConfig();
Expand Down Expand Up @@ -310,17 +315,6 @@ absl::StatusOr<std::unique_ptr<cel::Runtime>> 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<cel::Kind> param_kinds;
param_kinds.reserve(overload_config.parameters.size());
for (const cel::Config::TypeInfo& parameter :
Expand All @@ -333,15 +327,20 @@ absl::StatusOr<std::unique_ptr<cel::Runtime>> 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<PyCelFunctionAdapter>(
function_config.name, PyCelType::FromCelType(return_type),
std::move(py_function))));
descriptor, std::make_unique<PyCelFunctionAdapter>(
function_config.name,
PyCelType::FromCelType(return_type), it->second)));
}
}
return std::move(builder).Build();
Expand Down
2 changes: 1 addition & 1 deletion cel_expr_python/py_cel_env_internal.h
Original file line number Diff line number Diff line change
Expand Up @@ -76,7 +76,7 @@ class PyCelEnvInternal {
PyObject* py_descriptor_pool,
const std::unordered_map<std::string, PyCelType>& variable_types,
const std::vector<PyObject*>& extensions,
cel::ExpressionContainer container,
const cel::ExpressionContainer& container,
const std::vector<std::shared_ptr<PyCelFunctionDecl>>& functions,
const std::unordered_map<std::string, py::object>& function_impls);

Expand Down
37 changes: 25 additions & 12 deletions cel_expr_python/py_cel_expression.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand All @@ -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"
Expand Down Expand Up @@ -129,7 +129,8 @@ absl::StatusOr<PyCelExpression> 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_);
Expand All @@ -155,7 +156,10 @@ PyCelType PyCelExpression::GetReturnType() {
}

absl::StatusOr<const cel::Program*> 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();
}
Expand All @@ -179,17 +183,26 @@ absl::StatusOr<const cel::Program*> PyCelExpression::GetProgram() {
absl::StatusOr<PyCelValue> PyCelExpression::Eval(
const PyCelActivation& activation) {
ABSL_CHECK(PyGILState_Check());
CEL_PYTHON_ASSIGN_OR_RETURN(const cel::Program* program, GetProgram());
std::shared_ptr<PyCelArena> arena = activation.GetArena();
std::shared_ptr<PyCelEnvInternal> 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));
}

Expand Down
10 changes: 8 additions & 2 deletions cel_expr_python/py_cel_expression.h
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -70,7 +70,13 @@ class PyCelExpression {
std::variant<cel::expr::ParsedExpr, cel::expr::CheckedExpr>
expr_;
std::shared_ptr<PyCelEnvInternal> 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> cel_program_ ABSL_GUARDED_BY(mutex_);
};

Expand Down
Loading
Loading