diff --git a/src/auth/models.cpp b/src/auth/models.cpp index 08bb389c5..0910c92da 100644 --- a/src/auth/models.cpp +++ b/src/auth/models.cpp @@ -87,8 +87,6 @@ std::string PermissionToString(Permission permission) { return "MODULE_WRITE"; case Permission::WEBSOCKET: return "WEBSOCKET"; - case Permission::EDGE_TYPES: - return "EDGE_TYPES"; } } @@ -185,11 +183,13 @@ bool operator==(const Permissions &first, const Permissions &second) { bool operator!=(const Permissions &first, const Permissions &second) { return !(first == second); } -AccessPermissions::AccessPermissions(const std::unordered_set &grants, - const std::unordered_set &denies) +const std::string ASTERISK = "*"; + +FineGrainedAccessPermissions::FineGrainedAccessPermissions(const std::unordered_set &grants, + const std::unordered_set &denies) : grants_(grants), denies_(denies) {} -PermissionLevel AccessPermissions::Has(const std::string &permission) const { +PermissionLevel FineGrainedAccessPermissions::Has(const std::string &permission) const { if ((denies_.size() == 1 && denies_.find(ASTERISK) != denies_.end()) || denies_.find(permission) != denies_.end()) { return PermissionLevel::DENY; } @@ -201,7 +201,7 @@ PermissionLevel AccessPermissions::Has(const std::string &permission) const { return PermissionLevel::NEUTRAL; } -void AccessPermissions::Grant(const std::string &permission) { +void FineGrainedAccessPermissions::Grant(const std::string &permission) { if (permission == ASTERISK) { grants_.clear(); grants_.insert(permission); @@ -224,7 +224,7 @@ void AccessPermissions::Grant(const std::string &permission) { } } -void AccessPermissions::Revoke(const std::string &permission) { +void FineGrainedAccessPermissions::Revoke(const std::string &permission) { if (permission == ASTERISK) { grants_.clear(); denies_.clear(); @@ -244,7 +244,7 @@ void AccessPermissions::Revoke(const std::string &permission) { } } -void AccessPermissions::Deny(const std::string &permission) { +void FineGrainedAccessPermissions::Deny(const std::string &permission) { if (permission == ASTERISK) { denies_.clear(); denies_.insert(permission); @@ -267,18 +267,14 @@ void AccessPermissions::Deny(const std::string &permission) { } } -std::unordered_set AccessPermissions::GetGrants() const { return grants_; } - -std::unordered_set AccessPermissions::GetDenies() const { return denies_; } - -nlohmann::json AccessPermissions::Serialize() const { +nlohmann::json FineGrainedAccessPermissions::Serialize() const { nlohmann::json data = nlohmann::json::object(); data["grants"] = grants_; data["denies"] = denies_; return data; } -AccessPermissions AccessPermissions::Deserialize(const nlohmann::json &data) { +FineGrainedAccessPermissions FineGrainedAccessPermissions::Deserialize(const nlohmann::json &data) { if (!data.is_object()) { throw AuthException("Couldn't load permissions data!"); } @@ -286,23 +282,32 @@ AccessPermissions AccessPermissions::Deserialize(const nlohmann::json &data) { return {data["grants"], data["denies"]}; } -std::unordered_set AccessPermissions::grants() const { return grants_; } -std::unordered_set AccessPermissions::denies() const { return denies_; } +const std::unordered_set &FineGrainedAccessPermissions::grants() const { return grants_; } +const std::unordered_set &FineGrainedAccessPermissions::denies() const { return denies_; } -bool operator==(const AccessPermissions &first, const AccessPermissions &second) { +bool operator==(const FineGrainedAccessPermissions &first, const FineGrainedAccessPermissions &second) { return first.grants() == second.grants() && first.denies() == second.denies(); } -bool operator!=(const AccessPermissions &first, const AccessPermissions &second) { return !(first == second); } +bool operator!=(const FineGrainedAccessPermissions &first, const FineGrainedAccessPermissions &second) { + return !(first == second); +} Role::Role(const std::string &rolename) : rolename_(utils::ToLowerCase(rolename)) {} -Role::Role(const std::string &rolename, const Permissions &permissions, const AccessPermissions &edgeTypePermissions) - : rolename_(utils::ToLowerCase(rolename)), permissions_(permissions), edgeTypePermissions_(edgeTypePermissions) {} +Role::Role(const std::string &rolename, const Permissions &permissions, + const FineGrainedAccessPermissions &fine_grained_access_permissions) + : rolename_(utils::ToLowerCase(rolename)), + permissions_(permissions), + fine_grained_access_permissions_(fine_grained_access_permissions) {} const std::string &Role::rolename() const { return rolename_; } const Permissions &Role::permissions() const { return permissions_; } Permissions &Role::permissions() { return permissions_; } +const FineGrainedAccessPermissions &Role::fine_grained_access_permissions() const { + return fine_grained_access_permissions_; +} +FineGrainedAccessPermissions &Role::fine_grained_access_permissions() { return fine_grained_access_permissions_; } const AccessPermissions &Role::edgeTypePermissions() const { return edgeTypePermissions_; } AccessPermissions &Role::edgeTypePermissions() { return edgeTypePermissions_; } @@ -311,7 +316,7 @@ nlohmann::json Role::Serialize() const { nlohmann::json data = nlohmann::json::object(); data["rolename"] = rolename_; data["permissions"] = permissions_.Serialize(); - data["edgeTypePermissions"] = edgeTypePermissions_.Serialize(); + data["fine_grained_access_permissions"] = fine_grained_access_permissions_.Serialize(); return data; } @@ -319,12 +324,14 @@ Role Role::Deserialize(const nlohmann::json &data) { if (!data.is_object()) { throw AuthException("Couldn't load role data!"); } - if (!data["rolename"].is_string() || !data["permissions"].is_object()) { + if (!data["rolename"].is_string() || !data["permissions"].is_object() || + !data["fine_grained_access_permissions"].is_object()) { throw AuthException("Couldn't load role data!"); } auto permissions = Permissions::Deserialize(data["permissions"]); - auto edgeTypePermissions = AccessPermissions::Deserialize(data["edgeTypePermissions"]); - return {data["rolename"], permissions, edgeTypePermissions}; + auto fine_grained_access_permissions = + FineGrainedAccessPermissions::Deserialize(data["fine_grained_access_permissions"]); + return {data["rolename"], permissions, fine_grained_access_permissions}; } bool operator==(const Role &first, const Role &second) { @@ -335,11 +342,11 @@ bool operator==(const Role &first, const Role &second) { User::User(const std::string &username) : username_(utils::ToLowerCase(username)) {} User::User(const std::string &username, const std::string &password_hash, const Permissions &permissions, - const AccessPermissions &edgeTypePermissions) + const FineGrainedAccessPermissions &fine_grained_access_permissions) : username_(utils::ToLowerCase(username)), password_hash_(password_hash), permissions_(permissions), - edgeTypePermissions_(edgeTypePermissions) {} + fine_grained_access_permissions_(fine_grained_access_permissions) {} bool User::CheckPassword(const std::string &password) { if (password_hash_.empty()) return true; @@ -388,29 +395,35 @@ Permissions User::GetPermissions() const { return permissions_; } -AccessPermissions User::GetEdgeTypePermissions() const { +FineGrainedAccessPermissions User::GetFineGrainedAccessPermissions() const { if (role_) { std::unordered_set resultGrants; - std::set_union(edgeTypePermissions_.grants().begin(), edgeTypePermissions_.grants().end(), - role_->edgeTypePermissions().grants().begin(), role_->edgeTypePermissions().grants().end(), + std::set_union(fine_grained_access_permissions_.grants().begin(), fine_grained_access_permissions_.grants().end(), + role_->fine_grained_access_permissions().grants().begin(), + role_->fine_grained_access_permissions().grants().end(), std::inserter(resultGrants, resultGrants.begin())); std::unordered_set resultDenies; - std::set_union(edgeTypePermissions_.denies().begin(), edgeTypePermissions_.denies().end(), - role_->edgeTypePermissions().denies().begin(), role_->edgeTypePermissions().denies().end(), + std::set_union(fine_grained_access_permissions_.denies().begin(), fine_grained_access_permissions_.denies().end(), + role_->fine_grained_access_permissions().denies().begin(), + role_->fine_grained_access_permissions().denies().end(), std::inserter(resultDenies, resultDenies.begin())); return {resultGrants, resultDenies}; } - return edgeTypePermissions_; + return fine_grained_access_permissions_; } const std::string &User::username() const { return username_; } const Permissions &User::permissions() const { return permissions_; } Permissions &User::permissions() { return permissions_; } +const FineGrainedAccessPermissions &User::fine_grained_access_permissions() const { + return fine_grained_access_permissions_; +} +FineGrainedAccessPermissions &User::fine_grained_access_permissions() { return fine_grained_access_permissions_; } const AccessPermissions &User::edgeTypePermissions() const { return edgeTypePermissions_; } AccessPermissions &User::edgeTypePermissions() { return edgeTypePermissions_; } @@ -427,7 +440,7 @@ nlohmann::json User::Serialize() const { data["username"] = username_; data["password_hash"] = password_hash_; data["permissions"] = permissions_.Serialize(); - data["edgeTypePermissions"] = edgeTypePermissions_.Serialize(); + data["fine_grained_access_permissions"] = fine_grained_access_permissions_.Serialize(); // The role shouldn't be serialized here, it is stored as a foreign key. return data; } @@ -440,8 +453,9 @@ User User::Deserialize(const nlohmann::json &data) { throw AuthException("Couldn't load user data!"); } auto permissions = Permissions::Deserialize(data["permissions"]); - auto edgeTypePermissions = AccessPermissions::Deserialize(data["edgeTypePermissions"]); - return {data["username"], data["password_hash"], permissions, edgeTypePermissions}; + auto fine_grained_access_permissions = + FineGrainedAccessPermissions::Deserialize(data["fine_grained_access_permissions"]); + return {data["username"], data["password_hash"], permissions, fine_grained_access_permissions}; } bool operator==(const User &first, const User &second) { diff --git a/src/auth/models.hpp b/src/auth/models.hpp index a8290fe50..4079b64ee 100644 --- a/src/auth/models.hpp +++ b/src/auth/models.hpp @@ -10,6 +10,7 @@ #include #include +#include #include #include @@ -39,8 +40,7 @@ enum class Permission : uint64_t { STREAM = 1U << 17U, MODULE_READ = 1U << 18U, MODULE_WRITE = 1U << 19U, - WEBSOCKET = 1U << 20U, - EDGE_TYPES = 1U << 21U + WEBSOCKET = 1U << 20U }; // clang-format on @@ -90,10 +90,10 @@ bool operator==(const Permissions &first, const Permissions &second); bool operator!=(const Permissions &first, const Permissions &second); -class AccessPermissions final { +class FineGrainedAccessPermissions final { public: - AccessPermissions(const std::unordered_set &grants = {}, - const std::unordered_set &denies = {}); + FineGrainedAccessPermissions(const std::unordered_set &grants = {}, + const std::unordered_set &denies = {}); PermissionLevel Has(const std::string &permission) const; @@ -103,35 +103,35 @@ class AccessPermissions final { void Deny(const std::string &permission); - std::unordered_set GetGrants() const; - std::unordered_set GetDenies() const; - nlohmann::json Serialize() const; /// @throw AuthException if unable to deserialize. - static AccessPermissions Deserialize(const nlohmann::json &data); + static FineGrainedAccessPermissions Deserialize(const nlohmann::json &data); - std::unordered_set grants() const; - std::unordered_set denies() const; + const std::unordered_set &grants() const; + const std::unordered_set &denies() const; private: std::unordered_set grants_{}; std::unordered_set denies_{}; }; -bool operator==(const AccessPermissions &first, const AccessPermissions &second); +bool operator==(const FineGrainedAccessPermissions &first, const FineGrainedAccessPermissions &second); -bool operator!=(const AccessPermissions &first, const AccessPermissions &second); +bool operator!=(const FineGrainedAccessPermissions &first, const FineGrainedAccessPermissions &second); class Role final { public: Role(const std::string &rolename); - Role(const std::string &rolename, const Permissions &permissions, const AccessPermissions &edgeTypePermissions_); + Role(const std::string &rolename, const Permissions &permissions, + const FineGrainedAccessPermissions &fine_grained_access_permissions); const std::string &rolename() const; const Permissions &permissions() const; Permissions &permissions(); + const FineGrainedAccessPermissions &fine_grained_access_permissions() const; + FineGrainedAccessPermissions &fine_grained_access_permissions(); const AccessPermissions &edgeTypePermissions() const; AccessPermissions &edgeTypePermissions(); @@ -146,7 +146,7 @@ class Role final { private: std::string rolename_; Permissions permissions_; - AccessPermissions edgeTypePermissions_; + FineGrainedAccessPermissions fine_grained_access_permissions_; }; bool operator==(const Role &first, const Role &second); @@ -157,7 +157,7 @@ class User final { User(const std::string &username); User(const std::string &username, const std::string &password_hash, const Permissions &permissions, - const AccessPermissions &edgeTypePermissions_); + const FineGrainedAccessPermissions &fine_grained_access_permissions); /// @throw AuthException if unable to verify the password. bool CheckPassword(const std::string &password); @@ -170,12 +170,14 @@ class User final { void ClearRole(); Permissions GetPermissions() const; - AccessPermissions GetEdgeTypePermissions() const; + FineGrainedAccessPermissions GetFineGrainedAccessPermissions() const; const std::string &username() const; const Permissions &permissions() const; Permissions &permissions(); + const FineGrainedAccessPermissions &fine_grained_access_permissions() const; + FineGrainedAccessPermissions &fine_grained_access_permissions(); const AccessPermissions &edgeTypePermissions() const; AccessPermissions &edgeTypePermissions(); @@ -193,7 +195,7 @@ class User final { std::string username_; std::string password_hash_; Permissions permissions_; - AccessPermissions edgeTypePermissions_; + FineGrainedAccessPermissions fine_grained_access_permissions_; std::optional role_; }; diff --git a/src/glue/auth.cpp b/src/glue/auth.cpp index 512b98880..7f05d8045 100644 --- a/src/glue/auth.cpp +++ b/src/glue/auth.cpp @@ -57,8 +57,6 @@ auth::Permission PrivilegeToPermission(query::AuthQuery::Privilege privilege) { return auth::Permission::MODULE_WRITE; case query::AuthQuery::Privilege::WEBSOCKET: return auth::Permission::WEBSOCKET; - case query::AuthQuery::Privilege::EDGE_TYPES: - return auth::Permission::EDGE_TYPES; } } } // namespace memgraph::glue diff --git a/src/memgraph.cpp b/src/memgraph.cpp index 9fdee899d..86354c1f6 100644 --- a/src/memgraph.cpp +++ b/src/memgraph.cpp @@ -771,8 +771,8 @@ class AuthQueryHandler final : public memgraph::query::AuthQueryHandler { void GrantPrivilege(const std::string &user_or_role, const std::vector &privileges, - const std::vector &edgeTypes) override { - EditPermissions(user_or_role, privileges, edgeTypes, [](auto *permissions, const auto &permission) { + const std::vector &labels) override { + EditPermissions(user_or_role, privileges, labels, [](auto *permissions, const auto &permission) { // TODO (mferencevic): should we first check that the // privilege is granted/denied/revoked before // unconditionally granting/denying/revoking it? @@ -782,8 +782,8 @@ class AuthQueryHandler final : public memgraph::query::AuthQueryHandler { void DenyPrivilege(const std::string &user_or_role, const std::vector &privileges, - const std::vector &edgeTypes) override { - EditPermissions(user_or_role, privileges, edgeTypes, [](auto *permissions, const auto &permission) { + const std::vector &labels) override { + EditPermissions(user_or_role, privileges, labels, [](auto *permissions, const auto &permission) { // TODO (mferencevic): should we first check that the // privilege is granted/denied/revoked before // unconditionally granting/denying/revoking it? @@ -793,8 +793,8 @@ class AuthQueryHandler final : public memgraph::query::AuthQueryHandler { void RevokePrivilege(const std::string &user_or_role, const std::vector &privileges, - const std::vector &edgeTypes) override { - EditPermissions(user_or_role, privileges, edgeTypes, [](auto *permissions, const auto &permission) { + const std::vector &labels) override { + EditPermissions(user_or_role, privileges, labels, [](auto *permissions, const auto &permission) { // TODO (mferencevic): should we first check that the // privilege is granted/denied/revoked before // unconditionally granting/denying/revoking it? @@ -806,7 +806,7 @@ class AuthQueryHandler final : public memgraph::query::AuthQueryHandler { template void EditPermissions(const std::string &user_or_role, const std::vector &privileges, - const std::vector &edgeTypes, const TEditFun &edit_fun) { + const std::vector &labels, const TEditFun &edit_fun) { if (!std::regex_match(user_or_role, name_regex_)) { throw memgraph::query::QueryRuntimeException("Invalid user or role name."); } @@ -826,18 +826,18 @@ class AuthQueryHandler final : public memgraph::query::AuthQueryHandler { for (const auto &permission : permissions) { edit_fun(&user->permissions(), permission); } - - for (const auto &edgeType : edgeTypes) { - edit_fun(&user->edgeTypePermissions(), edgeType); + for (const auto &label : labels) { + edit_fun(&user->fine_grained_access_permissions(), label); } locked_auth->SaveUser(*user); } else { for (const auto &permission : permissions) { edit_fun(&role->permissions(), permission); } - for (const auto &edgeType : edgeTypes) { - edit_fun(&role->edgeTypePermissions(), edgeType); + for (const auto &label : labels) { + edit_fun(&user->fine_grained_access_permissions(), label); } + locked_auth->SaveRole(*role); } } catch (const memgraph::auth::AuthException &e) { diff --git a/src/query/context.hpp b/src/query/context.hpp index 606a5ecf6..cbcc01423 100644 --- a/src/query/context.hpp +++ b/src/query/context.hpp @@ -15,6 +15,7 @@ #include "query/access_checker.hpp" #include "query/common.hpp" +#include "query/fine_grained_access_checker.hpp" #include "query/frontend/semantic/symbol_table.hpp" #include "query/metadata.hpp" #include "query/parameters.hpp" @@ -73,7 +74,7 @@ struct ExecutionContext { ExecutionStats execution_stats; TriggerContextCollector *trigger_context_collector{nullptr}; utils::AsyncTimer timer; - AccessChecker *access_checker{nullptr}; + FineGrainedAccessChecker *fine_grained_access_checker{nullptr}; }; static_assert(std::is_move_assignable_v, "ExecutionContext must be move assignable!"); diff --git a/src/query/fine_grained_access_checker.hpp b/src/query/fine_grained_access_checker.hpp new file mode 100644 index 000000000..cd508734a --- /dev/null +++ b/src/query/fine_grained_access_checker.hpp @@ -0,0 +1,24 @@ +// Copyright 2022 Memgraph Ltd. +// +// Use of this software is governed by the Business Source License +// included in the file licenses/BSL.txt; by using this file, you agree to be bound by the terms of the Business Source +// License, and you may not use this file except in compliance with the Business Source License. +// +// As of the Change Date specified in that file, in accordance with +// the Business Source License, use of this software will be governed +// by the Apache License, Version 2.0, included in the file +// licenses/APL.txt. + +#pragma once + +#include "auth/models.hpp" +#include "query/frontend/ast/ast.hpp" +#include "storage/v2/id_types.hpp" + +namespace memgraph::query { +class FineGrainedAccessChecker { + public: + virtual bool IsUserAuthorizedLabels(const std::vector &label, + memgraph::query::DbAccessor *dba) const = 0; +}; +} // namespace memgraph::query diff --git a/src/query/frontend/ast/ast.lcp b/src/query/frontend/ast/ast.lcp index 63cce8a89..32a7cad0e 100644 --- a/src/query/frontend/ast/ast.lcp +++ b/src/query/frontend/ast/ast.lcp @@ -2234,6 +2234,7 @@ cpp<# (:serialize (:slk)) (:clone)) + (lcp:define-class auth-query (query) ((action "Action" :scope :public) (user "std::string" :scope :public) @@ -2242,8 +2243,9 @@ cpp<# (password "Expression *" :initval "nullptr" :scope :public :slk-save #'slk-save-ast-pointer :slk-load (slk-load-ast-pointer "Expression")) - (edgeTypes "std::vector" :scope :public) - (privileges "std::vector" :scope :public)) + (privileges "std::vector" :scope :public) + (labels "std::vector" :scope :public) + (edgeTypes "std::vector" :scope :public)) (:public (lcp:define-enum action (create-role drop-role show-roles create-user set-password drop-user @@ -2254,7 +2256,7 @@ cpp<# (lcp:define-enum privilege (create delete match merge set remove index stats auth constraint dump replication durability read_file free_memory trigger config stream module_read module_write - websocket edge_types) + websocket) (:serialize)) #>cpp AuthQuery() = default; @@ -2265,15 +2267,16 @@ cpp<# #>cpp AuthQuery(Action action, std::string user, std::string role, std::string user_or_role, Expression *password, - std::vector edgeTypes, - std::vector privileges) + std::vector privileges, + std::vector labels, std::vector edgeTypes) : action_(action), user_(user), role_(role), user_or_role_(user_or_role), password_(password), - edgetypes_(edgeTypes), - privileges_(privileges) {} + privileges_(privileges), + labels_(labels), + edgetypes_(edgeTypes) {} cpp<#) (:private #>cpp @@ -2298,7 +2301,7 @@ const std::vector kPrivilegesAll = { AuthQuery::Privilege::FREE_MEMORY, AuthQuery::Privilege::TRIGGER, AuthQuery::Privilege::CONFIG, AuthQuery::Privilege::STREAM, AuthQuery::Privilege::MODULE_READ, AuthQuery::Privilege::MODULE_WRITE, - AuthQuery::Privilege::WEBSOCKET, }; + AuthQuery::Privilege::WEBSOCKET}; cpp<# (lcp:define-class info-query (query) diff --git a/src/query/frontend/ast/cypher_main_visitor.cpp b/src/query/frontend/ast/cypher_main_visitor.cpp index 17781538c..8aa2462d9 100644 --- a/src/query/frontend/ast/cypher_main_visitor.cpp +++ b/src/query/frontend/ast/cypher_main_visitor.cpp @@ -1276,6 +1276,8 @@ antlrcpp::Any CypherMainVisitor::visitGrantPrivilege(MemgraphCypher::GrantPrivil for (auto *privilege : ctx->privilegeList()->privilege()) { if (privilege->EDGE_TYPES()) { auth->edgetypes_ = privilege->edgeTypeList()->accept(this).as>(); + } else if (privilege->LABELS()) { + auth->labels_ = privilege->labelList()->accept(this).as>(); } else { auth->privileges_.push_back(privilege->accept(this)); } @@ -1296,7 +1298,11 @@ antlrcpp::Any CypherMainVisitor::visitDenyPrivilege(MemgraphCypher::DenyPrivileg auth->user_or_role_ = ctx->userOrRole->accept(this).as(); if (ctx->privilegeList()) { for (auto *privilege : ctx->privilegeList()->privilege()) { - auth->privileges_.push_back(privilege->accept(this)); + if (privilege->LABELS()) { + auth->labels_ = privilege->labelList()->accept(this).as>(); + } else { + auth->privileges_.push_back(privilege->accept(this)); + } } } else { /* deny all privileges */ @@ -1314,7 +1320,11 @@ antlrcpp::Any CypherMainVisitor::visitRevokePrivilege(MemgraphCypher::RevokePriv auth->user_or_role_ = ctx->userOrRole->accept(this).as(); if (ctx->privilegeList()) { for (auto *privilege : ctx->privilegeList()->privilege()) { - auth->privileges_.push_back(privilege->accept(this)); + if (privilege->LABELS()) { + auth->labels_ = privilege->labelList()->accept(this).as>(); + } else { + auth->privileges_.push_back(privilege->accept(this)); + } } } else { /* revoke all privileges */ @@ -1323,6 +1333,19 @@ antlrcpp::Any CypherMainVisitor::visitRevokePrivilege(MemgraphCypher::RevokePriv return auth; } +antlrcpp::Any CypherMainVisitor::visitLabelList(MemgraphCypher::LabelListContext *ctx) { + std::vector labels; + if (ctx->listOfLabels()) { + for (auto *label : ctx->listOfLabels()->label()) { + labels.push_back(label->symbolicName()->accept(this).as()); + } + } else { + labels.emplace_back("*"); + } + + return labels; +} + /** * @return AuthQuery* */ diff --git a/src/query/frontend/ast/cypher_main_visitor.hpp b/src/query/frontend/ast/cypher_main_visitor.hpp index 14b4600c3..b954b1989 100644 --- a/src/query/frontend/ast/cypher_main_visitor.hpp +++ b/src/query/frontend/ast/cypher_main_visitor.hpp @@ -483,6 +483,11 @@ class CypherMainVisitor : public antlropencypher::MemgraphCypherBaseVisitor { */ antlrcpp::Any visitShowPrivileges(MemgraphCypher::ShowPrivilegesContext *ctx) override; + /** + * @return AuthQuery::LabelList + */ + antlrcpp::Any visitLabelList(MemgraphCypher::LabelListContext *ctx) override; + /** * @return AuthQuery* */ diff --git a/src/query/frontend/opencypher/grammar/MemgraphCypher.g4 b/src/query/frontend/opencypher/grammar/MemgraphCypher.g4 index f60a0f079..6ab820303 100644 --- a/src/query/frontend/opencypher/grammar/MemgraphCypher.g4 +++ b/src/query/frontend/opencypher/grammar/MemgraphCypher.g4 @@ -45,6 +45,7 @@ memgraphCypherKeyword : cypherKeyword | DENY | DROP | DUMP + | EDGE_TYPES | EXECUTE | FOR | FOREACH @@ -56,6 +57,7 @@ memgraphCypherKeyword : cypherKeyword | IDENTIFIED | ISOLATION | KAFKA + | LABELS | LEVEL | LOAD | LOCK @@ -255,6 +257,7 @@ privilege : CREATE | MODULE_WRITE | WEBSOCKET | EDGE_TYPES edgeTypes = edgeTypeList + | LABELS labels=labelList ; privilegeList : privilege ( ',' privilege )* ; @@ -265,6 +268,11 @@ edgeTypeList : '*' | listOfEdgeTypes ; listOfEdgeTypes : edgeType ( ',' edgeType )* ; edgeType : COLON symbolicName ; +labelList : '*' | listOfLabels ; + +listOfLabels : label ( ',' label )* ; + +label : COLON symbolicName ; showPrivileges : SHOW PRIVILEGES FOR userOrRole=userOrRoleName ; diff --git a/src/query/frontend/opencypher/grammar/MemgraphCypherLexer.g4 b/src/query/frontend/opencypher/grammar/MemgraphCypherLexer.g4 index 18f0a25a2..35b94d9db 100644 --- a/src/query/frontend/opencypher/grammar/MemgraphCypherLexer.g4 +++ b/src/query/frontend/opencypher/grammar/MemgraphCypherLexer.g4 @@ -66,6 +66,7 @@ IDENTIFIED : I D E N T I F I E D ; IGNORE : I G N O R E ; ISOLATION : I S O L A T I O N ; KAFKA : K A F K A ; +LABELS : L A B E L S ; LEVEL : L E V E L ; LOAD : L O A D ; LOCK : L O C K ; diff --git a/src/query/frontend/stripped_lexer_constants.hpp b/src/query/frontend/stripped_lexer_constants.hpp index 8d07904d8..b87f1893b 100644 --- a/src/query/frontend/stripped_lexer_constants.hpp +++ b/src/query/frontend/stripped_lexer_constants.hpp @@ -207,6 +207,8 @@ const trie::Trie kKeywords = {"union", "websocket", "foreach", "edge_types"}; +"labels" +}; // 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 b2067268f..e0af31ff0 100644 --- a/src/query/interpreter.cpp +++ b/src/query/interpreter.cpp @@ -30,6 +30,7 @@ #include "query/db_accessor.hpp" #include "query/dump.hpp" #include "query/exceptions.hpp" +#include "query/fine_grained_access_checker.hpp" #include "query/frontend/ast/ast.hpp" #include "query/frontend/ast/ast_visitor.hpp" #include "query/frontend/ast/cypher_main_visitor.hpp" @@ -260,31 +261,26 @@ class ReplQueryHandler final : public query::ReplicationQueryHandler { storage::Storage *db_; }; -class AccessChecker final : public memgraph::query::AccessChecker { +class FineGrainedAccessChecker final : public memgraph::query::FineGrainedAccessChecker { public: - explicit AccessChecker(memgraph::auth::User *user) : user_{user} {} + explicit FineGrainedAccessChecker(memgraph::auth::User *user) : user_{user} {} - bool IsUserAuthorizedEdgeTypes(const std::vector &edgeTypes, - memgraph::query::DbAccessor *dba) const final { - auto edgeTypePermissions = user_->GetEdgeTypePermissions(); + bool IsUserAuthorizedLabels(const std::vector &labels, + memgraph::query::DbAccessor *dba) const final { + auto labelPermissions = user_->GetFineGrainedAccessPermissions(); - return std::any_of(edgeTypes.begin(), edgeTypes.end(), [edgeTypePermissions, dba](const auto edgeType) { - return edgeTypePermissions.Has(dba->EdgeTypeToName(edgeType)) == memgraph::auth::PermissionLevel::GRANT; + return std::any_of(labels.begin(), labels.end(), [&labelPermissions, &dba](const auto label) { + return labelPermissions.Has(dba->LabelToName(label)) == memgraph::auth::PermissionLevel::GRANT; }); } - std::vector GetGrantedEdgeTypesId(memgraph::query::DbAccessor *dba) const final { - auto edgeTypePermissions = user_->GetEdgeTypePermissions().GetGrants(); - std::vector edgeTypeIds{}; - for (auto edgeTypePermission : edgeTypePermissions) edgeTypeIds.push_back(dba->NameToEdgeType(edgeTypePermission)); - - return edgeTypeIds; - } - private: 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 @@ -305,7 +301,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); + std::vector labels = auth_query->labels_; auto password = EvaluateOptionalExpression(auth_query->password_, &evaluator); Callback callback; @@ -318,11 +314,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: @@ -336,7 +332,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>(); @@ -411,20 +407,20 @@ Callback HandleAuthQuery(AuthQuery *auth_query, AuthQueryHandler *auth, const Pa }; return callback; case AuthQuery::Action::GRANT_PRIVILEGE: - callback.fn = [auth, user_or_role, privileges, edgeTypes] { - auth->GrantPrivilege(user_or_role, privileges, edgeTypes); + callback.fn = [auth, user_or_role, privileges, labels] { + auth->GrantPrivilege(user_or_role, privileges, labels); return std::vector>(); }; return callback; case AuthQuery::Action::DENY_PRIVILEGE: - callback.fn = [auth, user_or_role, privileges, edgeTypes] { - auth->DenyPrivilege(user_or_role, privileges, edgeTypes); + callback.fn = [auth, user_or_role, privileges, labels] { + auth->DenyPrivilege(user_or_role, privileges, labels); return std::vector>(); }; return callback; case AuthQuery::Action::REVOKE_PRIVILEGE: { - callback.fn = [auth, user_or_role, privileges, edgeTypes] { - auth->RevokePrivilege(user_or_role, privileges, edgeTypes); + callback.fn = [auth, user_or_role, privileges, labels] { + auth->RevokePrivilege(user_or_role, privileges, labels); return std::vector>(); }; return callback; @@ -968,7 +964,11 @@ PullPlan::PullPlan(const std::shared_ptr plan, const Parameters &par #ifdef MG_ENTERPRISE if (username.has_value()) { memgraph::auth::User *user = interpreter_context->auth->GetUser(*username); +<<<<<<< HEAD ctx_.access_checker = new AccessChecker{user}; +======= + ctx_.fine_grained_access_checker = new FineGrainedAccessChecker{user}; +>>>>>>> E129-MG-implement-label-based-authorization } #endif if (interpreter_context->config.execution_timeout_sec > 0) { @@ -1302,6 +1302,8 @@ PreparedQuery PrepareProfileQuery(ParsedQuery parsed_query, bool in_explicit_tra parsed_inner_query.stripped_query.hash(), std::move(parsed_inner_query.ast_storage), cypher_query, parsed_inner_query.parameters, parsed_inner_query.is_cacheable ? &interpreter_context->plan_cache : nullptr, dba); auto rw_type_checker = plan::ReadWriteTypeChecker(); + auto optional_username = StringPointerToOptional(username); + rw_type_checker.InferRWType(const_cast(cypher_query_plan->plan())); auto optional_username = StringPointerToOptional(username); diff --git a/src/query/interpreter.hpp b/src/query/interpreter.hpp index 81c9cfc3d..155f5688a 100644 --- a/src/query/interpreter.hpp +++ b/src/query/interpreter.hpp @@ -13,6 +13,7 @@ #include +#include "auth/models.hpp" #include "query/auth_checker.hpp" #include "query/config.hpp" #include "query/context.hpp" @@ -103,15 +104,15 @@ class AuthQueryHandler { /// @throw QueryRuntimeException if an error ocurred. virtual void GrantPrivilege(const std::string &user_or_role, const std::vector &privileges, - const std::vector &edgeTypes) = 0; + const std::vector &labels) = 0; /// @throw QueryRuntimeException if an error ocurred. virtual void DenyPrivilege(const std::string &user_or_role, const std::vector &privileges, - const std::vector &edgeTypes) = 0; + const std::vector &labels) = 0; /// @throw QueryRuntimeException if an error ocurred. virtual void RevokePrivilege(const std::string &user_or_role, const std::vector &privileges, - const std::vector &edgeTypes) = 0; + const std::vector &labels) = 0; }; enum class QueryHandlerResult { COMMIT, ABORT, NOTHING };