diff --git a/packages/react-native-executorch/cpp/extensions/nlp/tokenizer.cpp b/packages/react-native-executorch/cpp/extensions/nlp/tokenizer.cpp index e17109e86f..01ece72339 100644 --- a/packages/react-native-executorch/cpp/extensions/nlp/tokenizer.cpp +++ b/packages/react-native-executorch/cpp/extensions/nlp/tokenizer.cpp @@ -67,6 +67,18 @@ TokenizerHostObject::TokenizerHostObject(std::string tokenizerPath) } } +std::unique_lock TokenizerHostObject::tryLockUnique(jsi::Runtime &rt, + std::string_view context) { + std::unique_lock lock(mutex_, std::try_to_lock); + if (!lock.owns_lock()) { + throw jsi::JSError(rt, std::format("{} is currently in use", context)); + } + if (!tokenizer_) { + throw jsi::JSError(rt, std::format("{} has been disposed", context)); + } + return lock; +} + jsi::Value TokenizerHostObject::get(jsi::Runtime &rt, const jsi::PropNameID &name) { auto nameStr = name.utf8(rt); @@ -81,14 +93,7 @@ jsi::Value TokenizerHostObject::get(jsi::Runtime &rt, const jsi::PropNameID &nam throw jsi::JSError(rt, "encode: Usage: encode(text)"); } - std::unique_lock lock(self->mutex_, std::try_to_lock); - if (!lock.owns_lock()) { - throw jsi::JSError(rt, "encode: Tokenizer is currently in use"); - } - - if (!self->tokenizer_) { - throw jsi::JSError(rt, "encode: Tokenizer has been disposed"); - } + auto lock = self->tryLockUnique(rt, "encode: Tokenizer"); auto text = conversions::asType(rt, "encode: text", args[0]); auto tokens = unwrap(rt, "encode: Failed to encode input", @@ -112,14 +117,7 @@ jsi::Value TokenizerHostObject::get(jsi::Runtime &rt, const jsi::PropNameID &nam skipSpecialTokens = conversions::asType(rt, "decode: skipSpecialTokens", args[1]); } - std::unique_lock lock(self->mutex_, std::try_to_lock); - if (!lock.owns_lock()) { - throw jsi::JSError(rt, "decode: Tokenizer is currently in use"); - } - - if (!self->tokenizer_) { - throw jsi::JSError(rt, "decode: Tokenizer has been disposed"); - } + auto lock = self->tryLockUnique(rt, "decode: Tokenizer"); auto tokens = conversions::asVector(rt, "decode: tokens", args[0]); @@ -142,14 +140,7 @@ jsi::Value TokenizerHostObject::get(jsi::Runtime &rt, const jsi::PropNameID &nam throw jsi::JSError(rt, "getVocabSize: Usage: getVocabSize()"); } - std::unique_lock lock(self->mutex_, std::try_to_lock); - if (!lock.owns_lock()) { - throw jsi::JSError(rt, "getVocabSize: Tokenizer is currently in use"); - } - - if (!self->tokenizer_) { - throw jsi::JSError(rt, "getVocabSize: Tokenizer has been disposed"); - } + auto lock = self->tryLockUnique(rt, "getVocabSize: Tokenizer"); return static_cast(self->tokenizer_->vocab_size()); }; @@ -163,14 +154,7 @@ jsi::Value TokenizerHostObject::get(jsi::Runtime &rt, const jsi::PropNameID &nam throw jsi::JSError(rt, "idToToken: Usage: idToToken(id)"); } - std::unique_lock lock(self->mutex_, std::try_to_lock); - if (!lock.owns_lock()) { - throw jsi::JSError(rt, "idToToken: Tokenizer is currently in use"); - } - - if (!self->tokenizer_) { - throw jsi::JSError(rt, "idToToken: Tokenizer has been disposed"); - } + auto lock = self->tryLockUnique(rt, "idToToken: Tokenizer"); auto tokenId = conversions::asType(rt, "idToToken: id", args[0]); auto token = unwrap(rt, "idToToken: Failed to convert id to token", @@ -188,14 +172,7 @@ jsi::Value TokenizerHostObject::get(jsi::Runtime &rt, const jsi::PropNameID &nam throw jsi::JSError(rt, "tokenToId: Usage: tokenToId(token)"); } - std::unique_lock lock(self->mutex_, std::try_to_lock); - if (!lock.owns_lock()) { - throw jsi::JSError(rt, "tokenToId: Tokenizer is currently in use"); - } - - if (!self->tokenizer_) { - throw jsi::JSError(rt, "tokenToId: Tokenizer has been disposed"); - } + auto lock = self->tryLockUnique(rt, "tokenToId: Tokenizer"); auto token = conversions::asType(rt, "tokenToId: token", args[0]); auto tokenId = unwrap(rt, "tokenToId: Failed to convert token to id", diff --git a/packages/react-native-executorch/cpp/extensions/nlp/tokenizer.h b/packages/react-native-executorch/cpp/extensions/nlp/tokenizer.h index aac5aa2c0c..77a4b9d638 100644 --- a/packages/react-native-executorch/cpp/extensions/nlp/tokenizer.h +++ b/packages/react-native-executorch/cpp/extensions/nlp/tokenizer.h @@ -3,6 +3,7 @@ #include #include #include +#include #include #include @@ -20,6 +21,9 @@ class TokenizerHostObject : public facebook::jsi::HostObject, std::vector getPropertyNames(facebook::jsi::Runtime &rt) override; private: + [[nodiscard]] std::unique_lock tryLockUnique(facebook::jsi::Runtime &rt, + std::string_view context); + std::string tokenizerPath_; std::unique_ptr tokenizer_; std::mutex mutex_;