diff --git a/Source/ocdm/CMakeLists.txt b/Source/ocdm/CMakeLists.txt index 502e544f..ff0046d2 100644 --- a/Source/ocdm/CMakeLists.txt +++ b/Source/ocdm/CMakeLists.txt @@ -174,3 +174,12 @@ InstallCMakeConfig( InstallPackageConfig( TARGETS ${TARGET} DESCRIPTION "OCDM library") + +option(BUILD_OCDM_ENDPOINT_TESTS "Build OCDM endpoint tests" OFF) +if(BUILD_OCDM_ENDPOINT_TESTS) + enable_testing() + add_executable(ocdm_endpoint_test test_endpoint_validation.cpp) + target_include_directories(ocdm_endpoint_test PRIVATE ${CMAKE_CURRENT_LIST_DIR}) + set_target_properties(ocdm_endpoint_test PROPERTIES CXX_STANDARD ${CXX_STD} CXX_STANDARD_REQUIRED YES) + add_test(NAME ocdm_endpoint_test COMMAND ocdm_endpoint_test) +endif() diff --git a/Source/ocdm/EndpointValidation.h b/Source/ocdm/EndpointValidation.h new file mode 100644 index 00000000..47b07828 --- /dev/null +++ b/Source/ocdm/EndpointValidation.h @@ -0,0 +1,27 @@ +#pragma once + +#ifndef __WINDOWS__ +#include +#include +#endif + +#include + +namespace Thunder { +namespace OCDM { + +inline bool TrustedEndpoint(const std::string& endpoint) +{ +#ifdef __WINDOWS__ + return endpoint == "127.0.0.1:63000"; +#else + struct stat info; + if ((endpoint.empty()) || (endpoint[0] != '/') || (lstat(endpoint.c_str(), &info) != 0)) { + return false; + } + return S_ISSOCK(info.st_mode) && (info.st_uid == geteuid()) && ((info.st_mode & (S_IWGRP | S_IWOTH)) == 0); +#endif +} + +} +} diff --git a/Source/ocdm/open_cdm.cpp b/Source/ocdm/open_cdm.cpp index 710f5e9a..ba77d90c 100644 --- a/Source/ocdm/open_cdm.cpp +++ b/Source/ocdm/open_cdm.cpp @@ -19,6 +19,7 @@ #include "open_cdm.h" #include #include "open_cdm_impl.h" +#include "EndpointValidation.h" MODULE_NAME_DECLARATION(BUILD_REFERENCE) @@ -64,16 +65,21 @@ OpenCDMError StringToAllocatedBuffer(const std::string& source, char* destinatio TheOne(const TheOne&) = delete; TheOne& operator= (const TheOne&) = delete; - TheOne() { + TheOne() + : _valid(false) + { string connector; if ((Core::SystemInfo::GetEnvironment(_T("OPEN_CDM_SERVER"), connector) == false) || (connector.empty() == true)) { connector = _T("/tmp/ocdm"); } - Core::SingletonType::Create(connector.c_str()); + if (OCDM::TrustedEndpoint(connector)) { + Core::SingletonType::Create(connector.c_str()); + _valid = true; + } } ~TheOne() { - if( Core::SingletonType::Dispose() == true ) { + if (_valid && (Core::SingletonType::Dispose() == true)) { // if the accessor was disposed here because the destructor of the static instance was called there // was no proper dispose before (opencdm_dispose and/or Singleton::Dispose). // The static dispose might be incomplete or have side effects (e.g. Threads could already be killed) @@ -82,14 +88,16 @@ OpenCDMError StringToAllocatedBuffer(const std::string& source, char* destinatio } public: - OpenCDMAccessor& Instance() { - return (Core::SingletonType::Instance()); + OpenCDMAccessor* Instance() { + return (_valid ? &(Core::SingletonType::Instance()) : nullptr); } + private: + bool _valid; + } singleton; - OpenCDMAccessor& result = singleton.Instance(); - return &result; + return singleton.Instance(); } diff --git a/Source/ocdm/test_endpoint_validation.cpp b/Source/ocdm/test_endpoint_validation.cpp new file mode 100644 index 00000000..061ca5ef --- /dev/null +++ b/Source/ocdm/test_endpoint_validation.cpp @@ -0,0 +1,30 @@ +#include "EndpointValidation.h" + +#ifndef __WINDOWS__ +#include +#include +#include +#endif + +#include + +int main() +{ +#ifdef __WINDOWS__ + return Thunder::OCDM::TrustedEndpoint("127.0.0.1:63000") ? 0 : 1; +#else + const char path[] = "/tmp/ocdm-endpoint-test.sock"; + unlink(path); + const int descriptor = socket(AF_UNIX, SOCK_STREAM, 0); + sockaddr_un address {}; + address.sun_family = AF_UNIX; + std::strncpy(address.sun_path, path, sizeof(address.sun_path) - 1); + if ((descriptor < 0) || (bind(descriptor, reinterpret_cast(&address), sizeof(address)) != 0) || (chmod(path, 0600) != 0)) { + return 1; + } + const bool trusted = Thunder::OCDM::TrustedEndpoint(path); + close(descriptor); + unlink(path); + return trusted && !Thunder::OCDM::TrustedEndpoint("127.0.0.1:1") ? 0 : 2; +#endif +}