7#include <pluginlib/class_loader.hpp>
11#include <prompt_msgs/msg/prompt_history.hpp>
12#include <prompt_msgs/msg/prompt_transaction.hpp>
13#include <prompt_msgs/srv/embedding.hpp>
14#include <prompt_msgs/srv/prompt.hpp>
15#include <prompt_msgs/srv/tokenize.hpp>
16#include <rclcpp/rclcpp.hpp>
40 PromptBridge(
const rclcpp::NodeOptions& options = rclcpp::NodeOptions())
41 : Node(
"prompt_bridge", options)
48 if (shared_from_this())
53 catch (
const std::bad_weak_ptr&)
70 this->declare_parameter(
"frame_id",
"agent");
71 this->declare_parameter(
"cached_transactions", 10);
73 frame_id_ = this->get_parameter(
"frame_id").as_string();
80 this->declare_parameter(
"prompt_family_names", rclcpp::ParameterValue(std::vector<std::string>{}));
85 this->declare_parameter(
"prompt_family_plugins." + family_key, rclcpp::ParameterValue(
""));
86 std::string plugin = this->get_parameter(
"prompt_family_plugins." + family_key).as_string();
94 this->declare_parameter(
"embedding_family_names", rclcpp::ParameterValue(std::vector<std::string>{}));
99 this->declare_parameter(
"embedding_family_plugins." + family_key, rclcpp::ParameterValue(
""));
100 std::string plugin = this->get_parameter(
"embedding_family_plugins." + family_key).as_string();
108 this->declare_parameter(
"tokenizer_family_names", rclcpp::ParameterValue(std::vector<std::string>{}));
113 this->declare_parameter(
"tokenizer_family_plugins." + family_key, rclcpp::ParameterValue(
""));
114 std::string plugin = this->get_parameter(
"tokenizer_family_plugins." + family_key).as_string();
123 prompt_history_pub_ = this->create_publisher<prompt_msgs::msg::PromptHistory>(
"prompt/history", 1);
135 std::placeholders::_1, std::placeholders::_2));
136 RCLCPP_INFO(this->get_logger(),
"Prompt service created at 'prompt/prompt'");
140 std::placeholders::_1, std::placeholders::_2));
141 RCLCPP_INFO(this->get_logger(),
"Embedding service created at 'prompt/embedding'");
145 std::placeholders::_1, std::placeholders::_2));
146 RCLCPP_INFO(this->get_logger(),
"Tokenizer service created at 'prompt/tokenizer'");
148 RCLCPP_INFO(this->get_logger(),
"PromptBridge initialized");
159 uuid_generate_random(uuid);
161 uuid_unparse(uuid, uuid_str);
162 return std::string(uuid_str);
174 std::shared_ptr<prompt::BaseClass> prompt_provider_instance_;
180 prompt_provider_instance_->initialize(shared_from_this());
182 return prompt_provider_instance_;
186 RCLCPP_ERROR(this->get_logger(),
"Prompt family not found");
201 std::shared_ptr<prompt::BaseClass> embed_provider_instance_;
207 embed_provider_instance_->initialize(shared_from_this());
209 return embed_provider_instance_;
213 RCLCPP_ERROR(this->get_logger(),
"Embedding family not found");
228 std::shared_ptr<prompt::BaseClass> tokenizer_provider_instance_;
234 tokenizer_provider_instance_->initialize(shared_from_this());
236 return tokenizer_provider_instance_;
240 RCLCPP_ERROR(this->get_logger(),
"Tokenizer family not found");
257 void prompt_service_cb(
const std::shared_ptr<PromptSrv::Request> req, std::shared_ptr<PromptSrv::Response> res)
260 std::shared_ptr<prompt::BaseClass> prompt_provider_;
267 RCLCPP_ERROR(this->get_logger(),
"Failed to load model: %s", e.
what());
268 res->response.success =
false;
269 res->response.response =
"Failed to load model: " + std::string(e.
what());
272 catch (
const std::exception& e)
274 RCLCPP_ERROR(this->get_logger(),
"Unexpected error while loading model: %s", e.
what());
275 res->response.success =
false;
276 res->response.response =
"Unexpected error while loading model: " + std::string(e.
what());
288 auto pre_send_time = this->now();
291 if (req->prompt.use_chat_mode)
298 if (req->prompt.use_cache)
301 if (req->prompt.flush_cache)
304 RCLCPP_WARN(this->get_logger(),
"Flushing cache requested on new prompt with no uuid. Prompt will be "
305 "processed without caching.");
308 result = prompt_provider_->sendPrompt(input);
316 dialogue.
role =
"user";
317 dialogue.
content = req->prompt.prompt;
322 dialogue2.
role =
"assistant";
329 RCLCPP_INFO(this->get_logger(),
"New chat prompt processed. Generated UUID: %s", uuid.c_str());
337 dialogue.
role =
"user";
338 dialogue.
content = req->prompt.prompt;
347 RCLCPP_INFO(this->get_logger(),
"New chat prompt cached. Generated UUID: %s", uuid.c_str());
357 result = prompt_provider_->sendPrompt(input);
362 dialogue.
role =
"user";
363 dialogue.
content = req->prompt.prompt;
368 dialogue2.
role =
"assistant";
376 RCLCPP_INFO(this->get_logger(),
"New chat prompt processed. Generated UUID: %s", uuid.c_str());
382 if (req->prompt.use_cache)
384 if (req->prompt.flush_cache)
395 dialogue.
role =
"user";
396 dialogue.
content = req->prompt.prompt;
401 dialogue2.
role =
"assistant";
408 RCLCPP_INFO(this->get_logger(),
"Chat prompt processed with flushed cache. UUID: %s", uuid.c_str());
419 conv_it->second.back().content +=
" " + req->prompt.prompt;
426 dialogue.
role =
"user";
427 dialogue.
content = req->prompt.prompt;
436 RCLCPP_INFO(this->get_logger(),
"Chat prompt cached. UUID: %s", uuid.c_str());
451 dialogue.
role =
"user";
452 dialogue.
content = req->prompt.prompt;
457 dialogue2.
role =
"assistant";
465 RCLCPP_INFO(this->get_logger(),
"Chat prompt processed. UUID: %s", uuid.c_str());
477 if (req->prompt.use_cache)
479 if (req->prompt.flush_cache)
482 RCLCPP_WARN(this->get_logger(),
"Flushing cache requested on new prompt with no uuid. Prompt will be "
483 "processed without caching.");
486 result = prompt_provider_->sendPrompt(input);
491 RCLCPP_INFO(this->get_logger(),
"New chat prompt processed. No UUID generated due to flush request and "
492 "disabled chat mode");
500 dialogue.
role =
"user";
501 dialogue.
content = req->prompt.prompt;
511 RCLCPP_INFO(this->get_logger(),
"New prompt cached. Generated UUID: %s", uuid.c_str());
518 result = prompt_provider_->sendPrompt(input);
524 RCLCPP_INFO(this->get_logger(),
"New prompt processed.");
531 if (req->prompt.use_cache)
533 if (req->prompt.flush_cache)
542 RCLCPP_WARN(this->get_logger(),
"UUID provided for flush request does not exist in conversation history. "
543 "Processing prompt directly.");
544 result = prompt_provider_->sendPrompt(input);
560 RCLCPP_INFO(this->get_logger(),
"Prompt processed with flushed cache. UUID: %s discontinued.",
572 RCLCPP_WARN(this->get_logger(),
"uuid : %s does not exist in cache. Starting new cache", uuid.c_str());
574 dialogue.
role =
"user";
575 dialogue.
content = req->prompt.prompt;
590 RCLCPP_INFO(this->get_logger(),
"Prompt cached without flushing in non-chat mode. UUID: %s.", uuid.c_str());
597 RCLCPP_WARN(this->get_logger(),
"UUID provided for non-chat prompt with no caching. Ignoring UUID and "
598 "processing as a generic prompt");
600 result = prompt_provider_->sendPrompt(input);
606 RCLCPP_INFO(this->get_logger(),
"Prompt processed.");
615 RCLCPP_ERROR(this->get_logger(),
"Prompt request failed: %s", e.
what());
616 res->response.success =
false;
617 res->response.buffered =
false;
618 res->response.response = std::string(
"Prompt request failed: ") + e.
what();
621 catch (
const std::exception& e)
623 RCLCPP_ERROR(this->get_logger(),
"Unexpected prompt request failure: %s", e.
what());
624 res->response.success =
false;
625 res->response.buffered =
false;
626 res->response.response = std::string(
"Unexpected prompt request failure: ") + e.
what();
644 std::shared_ptr<EmbeddingSrv::Response> res)
647 std::shared_ptr<prompt::BaseClass> embedding_provider_;
654 RCLCPP_ERROR(this->get_logger(),
"Failed to load model: %s", e.
what());
655 res->output.success =
false;
656 res->output.error =
"Failed to load model: " + std::string(e.
what());
659 catch (
const std::exception& e)
661 RCLCPP_ERROR(this->get_logger(),
"Unexpected error while loading model: %s", e.
what());
662 res->output.success =
false;
663 res->output.error =
"Unexpected error while loading model: " + std::string(e.
what());
673 result = embedding_provider_->get_embeddings(input);
676 RCLCPP_INFO(this->get_logger(),
"Embedding request processed.");
680 RCLCPP_ERROR(this->get_logger(),
"Embedding request failed: %s", e.
what());
681 res->output.success =
false;
682 res->output.error = std::string(
"Embedding request failed: ") + e.
what();
684 catch (
const std::exception& e)
686 RCLCPP_ERROR(this->get_logger(),
"Unexpected embedding request failure: %s", e.
what());
687 res->output.success =
false;
688 res->output.error = std::string(
"Unexpected embedding request failure: ") + e.
what();
704 void tokenize_service_cb(
const std::shared_ptr<TokenizeSrv::Request> req, std::shared_ptr<TokenizeSrv::Response> res)
707 std::shared_ptr<prompt::BaseClass> tokenizer_provider_;
714 RCLCPP_ERROR(this->get_logger(),
"Failed to load model: %s", e.
what());
715 res->output.success =
false;
716 res->output.error =
"Failed to load model: " + std::string(e.
what());
719 catch (
const std::exception& e)
721 RCLCPP_ERROR(this->get_logger(),
"Unexpected error while loading model: %s", e.
what());
722 res->output.success =
false;
723 res->output.error =
"Unexpected error while loading model: " + std::string(e.
what());
733 result = tokenizer_provider_->get_tokens(input);
736 RCLCPP_INFO(this->get_logger(),
"Tokenization request processed.");
740 RCLCPP_ERROR(this->get_logger(),
"Tokenization request failed: %s", e.
what());
741 res->output.success =
false;
742 res->output.error = std::string(
"Tokenization request failed: ") + e.
what();
744 catch (
const std::exception& e)
746 RCLCPP_ERROR(this->get_logger(),
"Unexpected tokenization request failure: %s", e.
what());
747 res->output.success =
false;
748 res->output.error = std::string(
"Unexpected tokenization request failure: ") + e.
what();
774 rclcpp::Time prompt_time, rclcpp::Time response_time)
777 prompt_msgs::msg::PromptTransaction prompt_transaction = prompt_msgs::msg::PromptTransaction();
779 prompt_transaction.prompt.header = std_msgs::msg::Header();
780 prompt_transaction.prompt.header.stamp = prompt_time;
781 prompt_transaction.prompt.prompt =
prompt;
782 prompt_transaction.response.header = std_msgs::msg::Header();
783 prompt_transaction.response.header.stamp = response_time;
784 prompt_transaction.response.response = response;
PromptBridge.
Definition prompt_bridge.hpp:29
std::vector< std::string > prompt_families_keys_
Definition prompt_bridge.hpp:822
std::map< std::string, std::string > embedding_families_names_
Definition prompt_bridge.hpp:827
rclcpp::Publisher< prompt_msgs::msg::PromptHistory >::SharedPtr prompt_history_pub_
Definition prompt_bridge.hpp:811
rclcpp::Service< PromptSrv >::SharedPtr prompt_service_
Definition prompt_bridge.hpp:814
std::map< std::string, std::string > tokenizer_families_names_
Definition prompt_bridge.hpp:831
void history_timer()
Timer callback to publish prompt history at regular intervals.
Definition prompt_bridge.hpp:757
rclcpp::TimerBase::SharedPtr history_pub_timer_
Definition prompt_bridge.hpp:819
unsigned int transaction_limit_
Definition prompt_bridge.hpp:798
void embedding_service_cb(const std::shared_ptr< EmbeddingSrv::Request > req, std::shared_ptr< EmbeddingSrv::Response > res)
embedding service callback for embedding requests
Definition prompt_bridge.hpp:643
pluginlib::ClassLoader< prompt::BaseClass > tokenizer_loader
Definition prompt_bridge.hpp:808
prompt_msgs::msg::PromptHistory prompt_history_
Definition prompt_bridge.hpp:803
static std::string generate_uuid()
Generate a UUID string for prompt tracking.
Definition prompt_bridge.hpp:156
std::map< std::string, std::vector< prompt::PromptDialogue > > prompt_conversations_
Definition prompt_bridge.hpp:834
pluginlib::ClassLoader< prompt::BaseClass > prompt_loader
Definition prompt_bridge.hpp:806
rclcpp::Service< TokenizeSrv >::SharedPtr tokenizer_service_
Definition prompt_bridge.hpp:816
prompt_msgs::srv::Embedding EmbeddingSrv
Definition prompt_bridge.hpp:31
void prompt_service_cb(const std::shared_ptr< PromptSrv::Request > req, std::shared_ptr< PromptSrv::Response > res)
prompt service callback for instance prompts
Definition prompt_bridge.hpp:257
std::vector< std::string > embedding_families_keys_
Definition prompt_bridge.hpp:826
std::map< std::string, std::string > prompt_families_names_
Definition prompt_bridge.hpp:823
std::string provider_name_
Definition prompt_bridge.hpp:800
pluginlib::ClassLoader< prompt::BaseClass > embedding_loader
Definition prompt_bridge.hpp:807
prompt_msgs::srv::Prompt PromptSrv
Definition prompt_bridge.hpp:30
rclcpp::Service< EmbeddingSrv >::SharedPtr embedding_service_
Definition prompt_bridge.hpp:815
std::shared_ptr< prompt::BaseClass > load_prompt_model(std::string family)
Load a prompt model plugin based on the prompt family.
Definition prompt_bridge.hpp:172
prompt_msgs::srv::Tokenize TokenizeSrv
Definition prompt_bridge.hpp:32
std::shared_ptr< prompt::BaseClass > load_tokenizer_model(std::string family)
Load a tokenizer model plugin based on the tokenizer family.
Definition prompt_bridge.hpp:226
void initialize()
Initialize the PromptBridge node.
Definition prompt_bridge.hpp:64
std::vector< std::string > tokenizer_families_keys_
Definition prompt_bridge.hpp:830
std::string frame_id_
Definition prompt_bridge.hpp:797
void update_prompt_history(prompt_msgs::msg::Prompt prompt, prompt_msgs::msg::PromptResponse response, rclcpp::Time prompt_time, rclcpp::Time response_time)
Update the prompt history with a new prompt transaction.
Definition prompt_bridge.hpp:773
std::shared_ptr< prompt::BaseClass > load_embedding_model(std::string family)
Load an embedding model plugin based on the embedding family.
Definition prompt_bridge.hpp:199
void tokenize_service_cb(const std::shared_ptr< TokenizeSrv::Request > req, std::shared_ptr< TokenizeSrv::Response > res)
tokenizer service callback for tokenization requests
Definition prompt_bridge.hpp:704
PromptBridge(const rclcpp::NodeOptions &options=rclcpp::NodeOptions())
Construct a new Prompt Bridge object.
Definition prompt_bridge.hpp:40
Definition exceptions.hpp:10
virtual const char * what() const noexcept override
Definition exceptions.hpp:16
Definition base_class.hpp:8
static const prompt::PromptRequest fromMsg(const prompt_msgs::msg::Prompt &prompt)
Converts prompt_msgs::msg::Prompt into prompt::PromptRequest which is used internally in prompt tools...
Definition conversions.hpp:21
static const prompt_msgs::msg::PromptResponse toMsg(const prompt::PromptResponse &res)
Converts prompt::PromptResponse into prompt_msgs::msg::PromptResponse which is used in ros2 eco syste...
Definition conversions.hpp:90
Definition structs.hpp:46
Definition structs.hpp:61
Definition structs.hpp:17
std::string role
Definition structs.hpp:18
std::string content
Definition structs.hpp:19
Definition structs.hpp:24
Definition structs.hpp:35
bool buffered
Definition structs.hpp:37
std::string response
Definition structs.hpp:36
Definition structs.hpp:71
Definition structs.hpp:80