diff --git a/src/auth/models.hpp b/src/auth/models.hpp index df621c91c..a8290fe50 100644 --- a/src/auth/models.hpp +++ b/src/auth/models.hpp @@ -40,7 +40,7 @@ enum class Permission : uint64_t { MODULE_READ = 1U << 18U, MODULE_WRITE = 1U << 19U, WEBSOCKET = 1U << 20U, - EDGE_TYPES = 1U << 22U + EDGE_TYPES = 1U << 21U }; // clang-format on diff --git a/src/memgraph.cpp b/src/memgraph.cpp index e5395e0eb..9fdee899d 100644 --- a/src/memgraph.cpp +++ b/src/memgraph.cpp @@ -751,6 +751,24 @@ class AuthQueryHandler final : public memgraph::query::AuthQueryHandler { } } + memgraph::auth::User *GetUser(const std::string &username) override { + if (!std::regex_match(username, name_regex_)) { + throw memgraph::query::QueryRuntimeException("Invalid user name."); + } + try { + auto locked_auth = auth_->Lock(); + auto user = locked_auth->GetUser(username); + if (!user) { + throw memgraph::query::QueryRuntimeException("User '{}' doesn't exist .", username); + } + + return new memgraph::auth::User(*user); + + } catch (const memgraph::auth::AuthException &e) { + throw memgraph::query::QueryRuntimeException(e.what()); + } + } + void GrantPrivilege(const std::string &user_or_role, const std::vector &privileges, const std::vector &edgeTypes) override { diff --git a/src/query/context.hpp b/src/query/context.hpp index 12b1f0cac..606a5ecf6 100644 --- a/src/query/context.hpp +++ b/src/query/context.hpp @@ -13,6 +13,7 @@ #include +#include "query/access_checker.hpp" #include "query/common.hpp" #include "query/frontend/semantic/symbol_table.hpp" #include "query/metadata.hpp" @@ -72,6 +73,7 @@ struct ExecutionContext { ExecutionStats execution_stats; TriggerContextCollector *trigger_context_collector{nullptr}; utils::AsyncTimer timer; + AccessChecker *access_checker{nullptr}; }; static_assert(std::is_move_assignable_v, "ExecutionContext must be move assignable!"); diff --git a/src/query/frontend/opencypher/grammar/MemgraphCypherLexer.g4 b/src/query/frontend/opencypher/grammar/MemgraphCypherLexer.g4 index 55e5d53a2..18f0a25a2 100644 --- a/src/query/frontend/opencypher/grammar/MemgraphCypherLexer.g4 +++ b/src/query/frontend/opencypher/grammar/MemgraphCypherLexer.g4 @@ -114,3 +114,4 @@ USER : U S E R ; USERS : U S E R S ; VERSION : V E R S I O N ; WEBSOCKET : W E B S O C K E T ; +EDGE_TYPES : E D G E UNDERSCORE T Y P E S ; diff --git a/src/query/frontend/stripped_lexer_constants.hpp b/src/query/frontend/stripped_lexer_constants.hpp index 42b7b4aeb..8d07904d8 100644 --- a/src/query/frontend/stripped_lexer_constants.hpp +++ b/src/query/frontend/stripped_lexer_constants.hpp @@ -204,8 +204,9 @@ const trie::Trie kKeywords = {"union", "pulsar", "service_url", "version", - "websocket" - "foreach"}; + "websocket", + "foreach", + "edge_types"}; // Unicode codepoints that are allowed at the start of the unescaped name. const std::bitset kUnescapedNameAllowedStarts( diff --git a/src/query/interpreter.cpp b/src/query/interpreter.cpp index aa55d0fba..6a8e64711 100644 --- a/src/query/interpreter.cpp +++ b/src/query/interpreter.cpp @@ -44,9 +44,7 @@ #include "query/stream/common.hpp" #include "query/trigger.hpp" #include "query/typed_value.hpp" -#include "storage/v2/id_types.hpp" #include "storage/v2/property_value.hpp" -#include "storage/v2/replication/enums.hpp" #include "utils/algorithm.hpp" #include "utils/csv_parsing.hpp" #include "utils/event_counter.hpp" @@ -279,9 +277,6 @@ class AccessChecker final : public memgraph::query::AccessChecker { memgraph::auth::User *user_; }; -/// returns false if the replication role can't be set -/// @throw QueryRuntimeException if an error ocurred. - Callback HandleAuthQuery(AuthQuery *auth_query, AuthQueryHandler *auth, const Parameters ¶meters, DbAccessor *db_accessor) { // Empty frame for evaluation of password expression. This is OK since @@ -293,6 +288,7 @@ Callback HandleAuthQuery(AuthQuery *auth_query, AuthQueryHandler *auth, const Pa // TODO: MemoryResource for EvaluationContext, it should probably be passed as // the argument to Callback. evaluation_context.timestamp = QueryTimestamp(); + evaluation_context.parameters = parameters; ExpressionEvaluator evaluator(&frame, symbol_table, evaluation_context, db_accessor, storage::View::OLD); @@ -301,6 +297,7 @@ Callback HandleAuthQuery(AuthQuery *auth_query, AuthQueryHandler *auth, const Pa std::string user_or_role = auth_query->user_or_role_; std::vector privileges = auth_query->privileges_; std::vector edgeTypes = auth_query->edgetypes_; + // std::vector labels = NamesToLabels(labels, db_accessor); auto password = EvaluateOptionalExpression(auth_query->password_, &evaluator); Callback callback; @@ -313,10 +310,11 @@ Callback HandleAuthQuery(AuthQuery *auth_query, AuthQueryHandler *auth, const Pa AuthQuery::Action::REVOKE_PRIVILEGE, AuthQuery::Action::SHOW_PRIVILEGES, AuthQuery::Action::SHOW_USERS_FOR_ROLE, AuthQuery::Action::SHOW_ROLE_FOR_USER}; - if (license_check_result.HasError() && enterprise_only_methods.contains(auth_query->action_)) { - throw utils::BasicException( - utils::license::LicenseCheckErrorToString(license_check_result.GetError(), "advanced authentication features")); - } + // if (license_check_result.HasError() && enterprise_only_methods.contains(auth_query->action_)) { + // throw utils::BasicException( + // utils::license::LicenseCheckErrorToString(license_check_result.GetError(), "advanced authentication + // features")); + // } switch (auth_query->action_) { case AuthQuery::Action::CREATE_USER: @@ -330,7 +328,7 @@ Callback HandleAuthQuery(AuthQuery *auth_query, AuthQueryHandler *auth, const Pa // If the license is not valid we create users with admin access if (!valid_enterprise_license) { spdlog::warn("Granting all the privileges to {}.", username); - auth->GrantPrivilege(username, kPrivilegesAll, {"*"}); + auth->GrantPrivilege(username, kPrivilegesAll, {}); } return std::vector>(); @@ -918,7 +916,7 @@ struct PullPlanVector { struct PullPlan { explicit PullPlan(std::shared_ptr plan, const Parameters ¶meters, bool is_profile_query, DbAccessor *dba, InterpreterContext *interpreter_context, utils::MemoryResource *execution_memory, - TriggerContextCollector *trigger_context_collector = nullptr, + std::optional username, TriggerContextCollector *trigger_context_collector = nullptr, std::optional memory_limit = {}); std::optional Pull(AnyStream *stream, std::optional n, const std::vector &output_symbols, @@ -947,7 +945,8 @@ struct PullPlan { PullPlan::PullPlan(const std::shared_ptr plan, const Parameters ¶meters, const bool is_profile_query, DbAccessor *dba, InterpreterContext *interpreter_context, utils::MemoryResource *execution_memory, - TriggerContextCollector *trigger_context_collector, const std::optional memory_limit) + std::optional username, TriggerContextCollector *trigger_context_collector, + const std::optional memory_limit) : plan_(plan), cursor_(plan->plan().MakeCursor(execution_memory)), frame_(plan->symbol_table().max_position(), execution_memory), @@ -958,6 +957,12 @@ PullPlan::PullPlan(const std::shared_ptr plan, const Parameters &par ctx_.evaluation_context.parameters = parameters; ctx_.evaluation_context.properties = NamesToProperties(plan->ast_storage().properties_, dba); ctx_.evaluation_context.labels = NamesToLabels(plan->ast_storage().labels_, dba); +#ifdef MG_ENTERPRISE + if (username.has_value()) { + memgraph::auth::User *user = interpreter_context->auth->GetUser(*username); + ctx_.access_checker = new AccessChecker{user}; + } +#endif if (interpreter_context->config.execution_timeout_sec > 0) { ctx_.timer = utils::AsyncTimer{interpreter_context->config.execution_timeout_sec}; } @@ -1131,6 +1136,7 @@ PreparedQuery Interpreter::PrepareTransactionQuery(std::string_view query_upper) PreparedQuery PrepareCypherQuery(ParsedQuery parsed_query, std::map *summary, InterpreterContext *interpreter_context, DbAccessor *dba, utils::MemoryResource *execution_memory, std::vector *notifications, + const std::string *username, TriggerContextCollector *trigger_context_collector = nullptr) { auto *cypher_query = utils::Downcast(parsed_query.query); @@ -1139,6 +1145,7 @@ PreparedQuery PrepareCypherQuery(ParsedQuery parsed_query, std::mapmemory_limit_, cypher_query->memory_scale_); if (memory_limit) { @@ -1174,8 +1181,9 @@ PreparedQuery PrepareCypherQuery(ParsedQuery parsed_query, std::map(plan, parsed_query.parameters, false, dba, interpreter_context, - execution_memory, trigger_context_collector, memory_limit); + auto pull_plan = + std::make_shared(plan, parsed_query.parameters, false, dba, interpreter_context, execution_memory, + StringPointerToOptional(username), trigger_context_collector, memory_limit); return PreparedQuery{std::move(header), std::move(parsed_query.required_privileges), [pull_plan = std::move(pull_plan), output_symbols = std::move(output_symbols), summary]( AnyStream *stream, std::optional n) -> std::optional { @@ -1235,7 +1243,8 @@ PreparedQuery PrepareExplainQuery(ParsedQuery parsed_query, std::map *summary, InterpreterContext *interpreter_context, - DbAccessor *dba, utils::MemoryResource *execution_memory) { + DbAccessor *dba, utils::MemoryResource *execution_memory, + const std::string *username) { const std::string kProfileQueryStart = "profile "; MG_ASSERT(utils::StartsWith(utils::ToLowerCase(parsed_query.stripped_query.query()), kProfileQueryStart), @@ -1286,11 +1295,12 @@ PreparedQuery PrepareProfileQuery(ParsedQuery parsed_query, bool in_explicit_tra parsed_inner_query.parameters, parsed_inner_query.is_cacheable ? &interpreter_context->plan_cache : nullptr, dba); auto rw_type_checker = plan::ReadWriteTypeChecker(); rw_type_checker.InferRWType(const_cast(cypher_query_plan->plan())); + auto optional_username = StringPointerToOptional(username); return PreparedQuery{{"OPERATOR", "ACTUAL HITS", "RELATIVE TIME", "ABSOLUTE TIME"}, std::move(parsed_query.required_privileges), [plan = std::move(cypher_query_plan), parameters = std::move(parsed_inner_query.parameters), - summary, dba, interpreter_context, execution_memory, memory_limit, + summary, dba, interpreter_context, execution_memory, memory_limit, optional_username, // We want to execute the query we are profiling lazily, so we delay // the construction of the corresponding context. stats_and_total_time = std::optional{}, @@ -1299,7 +1309,7 @@ PreparedQuery PrepareProfileQuery(ParsedQuery parsed_query, bool in_explicit_tra // No output symbols are given so that nothing is streamed. if (!stats_and_total_time) { stats_and_total_time = PullPlan(plan, parameters, true, dba, interpreter_context, - execution_memory, nullptr, memory_limit) + execution_memory, optional_username, nullptr, memory_limit) .Pull(stream, {}, {}, summary); pull_plan = std::make_shared(ProfilingStatsToTable(*stats_and_total_time)); } @@ -1434,7 +1444,7 @@ PreparedQuery PrepareIndexQuery(ParsedQuery parsed_query, bool in_explicit_trans PreparedQuery PrepareAuthQuery(ParsedQuery parsed_query, bool in_explicit_transaction, std::map *summary, InterpreterContext *interpreter_context, - DbAccessor *dba, utils::MemoryResource *execution_memory) { + DbAccessor *dba, utils::MemoryResource *execution_memory, const std::string *username) { if (in_explicit_transaction) { throw UserModificationInMulticommandTxException(); } @@ -1454,8 +1464,8 @@ PreparedQuery PrepareAuthQuery(ParsedQuery parsed_query, bool in_explicit_transa [fn = callback.fn](Frame *, ExecutionContext *) { return fn(); }), 0.0, AstStorage{}, symbol_table)); - auto pull_plan = - std::make_shared(plan, parsed_query.parameters, false, dba, interpreter_context, execution_memory); + auto pull_plan = std::make_shared(plan, parsed_query.parameters, false, dba, interpreter_context, + execution_memory, StringPointerToOptional(username)); return PreparedQuery{ callback.header, std::move(parsed_query.required_privileges), [pull_plan = std::move(pull_plan), callback = std::move(callback), output_symbols = std::move(output_symbols), @@ -2168,7 +2178,7 @@ Interpreter::PrepareResult Interpreter::Prepare(const std::string &query_string, if (utils::Downcast(parsed_query.query)) { prepared_query = PrepareCypherQuery(std::move(parsed_query), &query_execution->summary, interpreter_context_, &*execution_db_accessor_, &query_execution->execution_memory, - &query_execution->notifications, + &query_execution->notifications, username, trigger_context_collector_ ? &*trigger_context_collector_ : nullptr); } else if (utils::Downcast(parsed_query.query)) { prepared_query = PrepareExplainQuery(std::move(parsed_query), &query_execution->summary, interpreter_context_, @@ -2176,7 +2186,7 @@ Interpreter::PrepareResult Interpreter::Prepare(const std::string &query_string, } else if (utils::Downcast(parsed_query.query)) { prepared_query = PrepareProfileQuery(std::move(parsed_query), in_explicit_transaction_, &query_execution->summary, interpreter_context_, &*execution_db_accessor_, - &query_execution->execution_memory_with_exception); + &query_execution->execution_memory_with_exception, username); } else if (utils::Downcast(parsed_query.query)) { prepared_query = PrepareDumpQuery(std::move(parsed_query), &query_execution->summary, &*execution_db_accessor_, &query_execution->execution_memory); @@ -2186,7 +2196,7 @@ Interpreter::PrepareResult Interpreter::Prepare(const std::string &query_string, } else if (utils::Downcast(parsed_query.query)) { prepared_query = PrepareAuthQuery(std::move(parsed_query), in_explicit_transaction_, &query_execution->summary, interpreter_context_, &*execution_db_accessor_, - &query_execution->execution_memory_with_exception); + &query_execution->execution_memory_with_exception, username); } else if (utils::Downcast(parsed_query.query)) { prepared_query = PrepareInfoQuery(std::move(parsed_query), in_explicit_transaction_, &query_execution->summary, interpreter_context_, interpreter_context_->db, diff --git a/src/query/interpreter.hpp b/src/query/interpreter.hpp index f04b0075c..81c9cfc3d 100644 --- a/src/query/interpreter.hpp +++ b/src/query/interpreter.hpp @@ -98,6 +98,9 @@ class AuthQueryHandler { virtual std::vector> GetPrivileges(const std::string &user_or_role) = 0; + /// @throw QueryRuntimeException if an error ocurred. + virtual memgraph::auth::User *GetUser(const std::string &username) = 0; + /// @throw QueryRuntimeException if an error ocurred. virtual void GrantPrivilege(const std::string &user_or_role, const std::vector &privileges, const std::vector &edgeTypes) = 0;