winget-cli

Unnamed repository; edit this file 'description' to name the repository.
Log | Files | Refs | README | LICENSE

RestSource.cpp (22586B)


      1 // Copyright (c) Microsoft Corporation.
      2 // Licensed under the MIT License.
      3 #include "pch.h"
      4 #include "RestSource.h"
      5 
      6 using namespace AppInstaller::Utility;
      7 
      8 namespace AppInstaller::Repository::Rest
      9 {
     10     namespace
     11     {
     12         using namespace AppInstaller::Repository::Rest::Schema;
     13 
     14         // The source reference used by package objects.
     15         struct SourceReference
     16         {
     17             SourceReference(const std::shared_ptr<RestSource>& source) :
     18                 m_source(source) {}
     19 
     20         protected:
     21             std::shared_ptr<RestSource> GetReferenceSource() const
     22             {
     23                 std::shared_ptr<RestSource> source = m_source.lock();
     24                 THROW_HR_IF(E_NOT_VALID_STATE, !source);
     25                 return source;
     26             }
     27 
     28         private:
     29             std::weak_ptr<RestSource> m_source;
     30         };
     31 
     32         // The IPackage implementation for Available packages from RestSource.
     33         struct RestPackage : public std::enable_shared_from_this<RestPackage>, public SourceReference, public IPackage, public ICompositePackage
     34         {
     35             static constexpr IPackageType PackageType = IPackageType::RestPackage;
     36 
     37             RestPackage(const std::shared_ptr<RestSource>& source, IRestClient::Package&& package) :
     38                 SourceReference(source), m_package(std::move(package))
     39             {
     40                 SortVersionsInternal();
     41             }
     42 
     43             // Inherited via IPackage
     44             Utility::LocIndString GetProperty(PackageProperty property) const override
     45             {
     46                 switch (property)
     47                 {
     48                 case PackageProperty::Id:
     49                     return Utility::LocIndString{ m_package.PackageInformation.PackageIdentifier };
     50                 case PackageProperty::Name:
     51                     return Utility::LocIndString{ m_package.PackageInformation.PackageName };
     52                 default:
     53                     THROW_HR(E_UNEXPECTED);
     54                 }
     55             }
     56 
     57             std::vector<Utility::LocIndString> GetMultiProperty(PackageMultiProperty property) const override;
     58 
     59             std::vector<PackageVersionKey> GetVersionKeys() const override
     60             {
     61                 std::shared_ptr<const RestSource> source = GetReferenceSource();
     62                 std::scoped_lock versionsLock{ m_packageVersionsLock };
     63 
     64                 std::vector<PackageVersionKey> result;
     65                 for (const auto& versionInfo : m_package.Versions)
     66                 {
     67                     result.emplace_back(
     68                         source->GetIdentifier(), versionInfo.VersionAndChannel.GetVersion().ToString(), versionInfo.VersionAndChannel.GetChannel().ToString());
     69                 }
     70 
     71                 return result;
     72             }
     73 
     74             std::shared_ptr<IPackageVersion> GetLatestVersion() const override
     75             {
     76                 std::scoped_lock versionsLock{ m_packageVersionsLock };
     77                 return GetLatestVersionInternal();
     78             }
     79 
     80             std::shared_ptr<IPackageVersion> GetVersion(const PackageVersionKey& versionKey) const override;
     81 
     82             Source GetSource() const override
     83             {
     84                 return Source{ GetReferenceSource() };
     85             }
     86 
     87             bool IsSame(const IPackage* other) const override
     88             {
     89                 const RestPackage* otherPackage = PackageCast<const RestPackage*>(other);
     90 
     91                 if (otherPackage)
     92                 {
     93                     return GetReferenceSource()->IsSame(otherPackage->GetReferenceSource().get()) &&
     94                         Utility::CaseInsensitiveEquals(m_package.PackageInformation.PackageIdentifier, otherPackage->m_package.PackageInformation.PackageIdentifier);
     95                 }
     96 
     97                 return false;
     98             }
     99 
    100             const void* CastTo(IPackageType type) const override
    101             {
    102                 if (type == PackageType)
    103                 {
    104                     return this;
    105                 }
    106 
    107                 return nullptr;
    108             }
    109 
    110             // Inherited via ICompositePackage
    111             std::shared_ptr<IPackage> GetInstalled() override
    112             {
    113                 return {};
    114             }
    115 
    116             std::vector<std::shared_ptr<IPackage>> GetAvailable() override
    117             {
    118                 return std::vector<std::shared_ptr<IPackage>>{ shared_from_this() };
    119             }
    120 
    121             // Helpers for PackageVersion interop
    122             const IRestClient::PackageInfo& PackageInfo() const
    123             {
    124                 return m_package.PackageInformation;
    125             }
    126 
    127             // This function is designed to handle the case where the only version that is returned by the
    128             // initial search is Unknown. In that case, we perform a search intended to trigger the optimized
    129             // path and directly get all manifests.
    130             bool HandleSingleUnknownVersion(IRestClient::VersionInfo& versionInfo)
    131             {
    132                 // If the calling version is unknown then we want to update it if we already
    133                 // have the results in the package.
    134                 if (versionInfo.VersionAndChannel.GetVersion().IsUnknown() && !versionInfo.Manifest)
    135                 {
    136                     std::scoped_lock versionsLock{ m_packageVersionsLock };
    137                     if (m_package.Versions.size() == 1 && m_package.Versions[0].VersionAndChannel.GetVersion().IsUnknown() && !m_package.Versions[0].Manifest)
    138                     {
    139                         SearchRequest request;
    140                         request.Filters.emplace_back(PackageMatchField::Id, MatchType::CaseInsensitive, m_package.PackageInformation.PackageIdentifier);
    141 
    142                         IRestClient::SearchResult result = GetReferenceSource()->GetRestClient().Search(request);
    143 
    144                         if (result.Matches.size() == 1)
    145                         {
    146                             m_package.Versions = std::move(result.Matches[0].Versions);
    147                             SortVersionsInternal();
    148                         }
    149                         else
    150                         {
    151                             // Unexpected, but just leave things as they are
    152                             AICLI_LOG(Repo, Warning, << "Found " << result.Matches.size() << " matches for optimized search of " << m_package.PackageInformation.PackageIdentifier);
    153                         }
    154                     }
    155 
    156                     if (!m_package.Versions.empty())
    157                     {
    158                         // The results are now sorted; either take the last one if it is unknown
    159                         // or the first one if it is not (aka latest).
    160                         if (m_package.Versions.back().VersionAndChannel.GetVersion().IsUnknown())
    161                         {
    162                             versionInfo = m_package.Versions.back();
    163                         }
    164                         else
    165                         {
    166                             versionInfo = m_package.Versions.front();
    167                         }
    168                     }
    169 
    170                     return true;
    171                 }
    172 
    173                 return false;
    174             }
    175 
    176         private:
    177             std::shared_ptr<RestPackage> NonConstSharedFromThis() const
    178             {
    179                 return const_cast<RestPackage*>(this)->shared_from_this();
    180             }
    181 
    182             // Must hold m_packageVersionsLock while calling this
    183             std::shared_ptr<IPackageVersion> GetLatestVersionInternal() const;
    184 
    185             // Must hold m_packageVersionsLock while calling this
    186             void SortVersionsInternal()
    187             {
    188                 std::sort(m_package.Versions.begin(), m_package.Versions.end(),
    189                     [](const IRestClient::VersionInfo& a, const IRestClient::VersionInfo& b)
    190                     {
    191                         return a.VersionAndChannel < b.VersionAndChannel;
    192                     });
    193             }
    194 
    195             IRestClient::Package m_package;
    196             // Protects access to m_package.Versions
    197             mutable std::mutex m_packageVersionsLock;
    198         };
    199 
    200         void GetMultiPropertyValues(
    201             const RestPackage* package,
    202             const IRestClient::VersionInfo& versionInfo,
    203             PackageVersionMultiProperty property,
    204             std::vector<Utility::LocIndString>& result,
    205             void (*Action)(std::vector<Utility::LocIndString>&, Utility::LocIndString&&))
    206         {
    207             switch (property)
    208             {
    209             case PackageVersionMultiProperty::PackageFamilyName:
    210                 for (const std::string& pfn : versionInfo.PackageFamilyNames)
    211                 {
    212                     Action(result, Utility::LocIndString{ pfn });
    213                 }
    214                 break;
    215             case PackageVersionMultiProperty::ProductCode:
    216                 for (const std::string& productCode : versionInfo.ProductCodes)
    217                 {
    218                     Action(result, Utility::LocIndString{ productCode });
    219                 }
    220                 break;
    221             case PackageVersionMultiProperty::UpgradeCode:
    222                 for (const std::string& upgradeCode : versionInfo.UpgradeCodes)
    223                 {
    224                     Action(result, Utility::LocIndString{ upgradeCode });
    225                 }
    226                 break;
    227             case PackageVersionMultiProperty::Name:
    228                 if (versionInfo.Manifest)
    229                 {
    230                     for (auto&& name : versionInfo.Manifest->GetPackageNames())
    231                     {
    232                         Action(result, Utility::LocIndString{ std::move(name) });
    233                     }
    234                 }
    235                 else
    236                 {
    237                     Action(result, Utility::LocIndString{ package->PackageInfo().PackageName });
    238                 }
    239                 break;
    240             case PackageVersionMultiProperty::Publisher:
    241                 if (versionInfo.Manifest)
    242                 {
    243                     for (auto&& publisher : versionInfo.Manifest->GetPublishers())
    244                     {
    245                         Action(result, Utility::LocIndString{ std::move(publisher) });
    246                     }
    247                 }
    248                 else
    249                 {
    250                     Action(result, Utility::LocIndString{ package->PackageInfo().Publisher });
    251                 }
    252                 break;
    253             case PackageVersionMultiProperty::Locale:
    254                 if (versionInfo.Manifest)
    255                 {
    256                     Action(result, Utility::LocIndString{ versionInfo.Manifest->DefaultLocalization.Locale });
    257                     for (const auto& loc : versionInfo.Manifest->Localizations)
    258                     {
    259                         Action(result, Utility::LocIndString{ loc.Locale });
    260                     }
    261                 }
    262                 break;
    263             }
    264         }
    265 
    266         std::vector<Utility::LocIndString> RestPackage::GetMultiProperty(PackageMultiProperty property) const
    267         {
    268             std::scoped_lock versionsLock{ m_packageVersionsLock };
    269             std::vector<Utility::LocIndString> result;
    270             PackageVersionMultiProperty mappedProperty = PackageMultiPropertyToPackageVersionMultiProperty(property);
    271 
    272             for (const auto& versionInfo : m_package.Versions)
    273             {
    274                 GetMultiPropertyValues(
    275                     this,
    276                     versionInfo,
    277                     mappedProperty,
    278                     result,
    279                     [](std::vector<Utility::LocIndString>& result, Utility::LocIndString&& string)
    280                     {
    281                         auto itr = std::lower_bound(result.begin(), result.end(), string);
    282 
    283                         if (itr == result.end() || *itr != string)
    284                         {
    285                             result.emplace(itr, std::move(string));
    286                         }
    287                     });
    288             }
    289 
    290             return result;
    291         }
    292 
    293         // The IPackageVersion impl for RestSource.
    294         struct PackageVersion : public SourceReference, public IPackageVersion
    295         {
    296             PackageVersion(
    297                 const std::shared_ptr<RestSource>& source, std::shared_ptr<RestPackage>&& package, IRestClient::VersionInfo versionInfo)
    298                 : SourceReference(source), m_package(std::move(package)), m_versionInfo(std::move(versionInfo)) {}
    299 
    300             // Inherited via IPackageVersion
    301             Utility::LocIndString GetProperty(PackageVersionProperty property) const override
    302             {
    303                 switch (property)
    304                 {
    305                 case PackageVersionProperty::SourceIdentifier:
    306                     return Utility::LocIndString{ GetReferenceSource()->GetIdentifier() };
    307                 case PackageVersionProperty::SourceName:
    308                     return Utility::LocIndString{ GetReferenceSource()->GetDetails().Name };
    309                 case PackageVersionProperty::Id:
    310                     return Utility::LocIndString{ m_package->PackageInfo().PackageIdentifier };
    311                 case PackageVersionProperty::Name:
    312                     return Utility::LocIndString{ m_package->PackageInfo().PackageName };
    313                 case PackageVersionProperty::Version:
    314                     return Utility::LocIndString{ m_versionInfo.VersionAndChannel.GetVersion().ToString() };
    315                 case PackageVersionProperty::Channel:
    316                     return Utility::LocIndString{ m_versionInfo.VersionAndChannel.GetChannel().ToString() };
    317                 case PackageVersionProperty::Publisher:
    318                     return Utility::LocIndString{ m_package->PackageInfo().Publisher };
    319                 case PackageVersionProperty::ArpMinVersion:
    320                     if (!m_versionInfo.ArpVersions.empty())
    321                     {
    322                         return Utility::LocIndString{ m_versionInfo.ArpVersions.front().ToString() };
    323                     }
    324                     else if (m_versionInfo.Manifest)
    325                     {
    326                         auto arpVersionRange = m_versionInfo.Manifest->GetArpVersionRange();
    327                         return arpVersionRange.IsEmpty() ? Utility::LocIndString{} : Utility::LocIndString{ arpVersionRange.GetMinVersion().ToString() };
    328                     }
    329                     else
    330                     {
    331                         return {};
    332                     }
    333                 case PackageVersionProperty::ArpMaxVersion:
    334                     if (!m_versionInfo.ArpVersions.empty())
    335                     {
    336                         return Utility::LocIndString{ m_versionInfo.ArpVersions.back().ToString() };
    337                     }
    338                     else if (m_versionInfo.Manifest)
    339                     {
    340                         auto arpVersionRange = m_versionInfo.Manifest->GetArpVersionRange();
    341                         return arpVersionRange.IsEmpty() ? Utility::LocIndString{} : Utility::LocIndString{ arpVersionRange.GetMaxVersion().ToString() };
    342                     }
    343                     else
    344                     {
    345                         return {};
    346                     }
    347                 default:
    348                     return {};
    349                 }
    350             }
    351 
    352             std::vector<Utility::LocIndString> GetMultiProperty(PackageVersionMultiProperty property) const override
    353             {
    354                 std::vector<Utility::LocIndString> result;
    355 
    356                 GetMultiPropertyValues(
    357                     m_package.get(),
    358                     m_versionInfo,
    359                     property,
    360                     result,
    361                     [](std::vector<Utility::LocIndString>& result, Utility::LocIndString&& string)
    362                     {
    363                         result.emplace_back(std::move(string));
    364                     });
    365 
    366                 return result;
    367             }
    368 
    369             Manifest::Manifest GetManifest() override
    370             {
    371                 AICLI_LOG(Repo, Verbose, << "Getting manifest");
    372 
    373                 if (m_versionInfo.Manifest)
    374                 {
    375                     return m_versionInfo.Manifest.value();
    376                 }
    377 
    378                 if (m_package->HandleSingleUnknownVersion(m_versionInfo) &&
    379                     m_versionInfo.Manifest)
    380                 {
    381                     return m_versionInfo.Manifest.value();
    382                 }
    383 
    384                 std::optional<Manifest::Manifest> manifest = GetReferenceSource()->GetRestClient().GetManifestByVersion(
    385                     m_package->PackageInfo().PackageIdentifier, m_versionInfo.VersionAndChannel.GetVersion().ToString(), m_versionInfo.VersionAndChannel.GetChannel().ToString());
    386 
    387                 if (!manifest)
    388                 {
    389                     AICLI_LOG(Repo, Verbose, << "Valid manifest not found for package: " << m_package->PackageInfo().PackageIdentifier);
    390                     return {};
    391                 }
    392                 
    393                 m_versionInfo.Manifest = std::move(manifest.value());
    394                 return m_versionInfo.Manifest.value();
    395             }
    396 
    397             Source GetSource() const override
    398             {
    399                 return Source{ GetReferenceSource() };
    400             }
    401 
    402             IPackageVersion::Metadata GetMetadata() const override
    403             {
    404                 IPackageVersion::Metadata result;
    405                 return result;
    406             }
    407 
    408         private:
    409             std::shared_ptr<RestPackage> m_package;
    410             IRestClient::VersionInfo m_versionInfo;
    411         };
    412 
    413         std::shared_ptr<IPackageVersion> RestPackage::GetVersion(const PackageVersionKey& versionKey) const
    414         {
    415             std::shared_ptr<RestSource> source = GetReferenceSource();
    416             std::scoped_lock versionsLock{ m_packageVersionsLock };
    417 
    418             // Ensure that this key targets this (or any) source
    419             if (!versionKey.SourceId.empty() && versionKey.SourceId != source->GetIdentifier())
    420             {
    421                 return {};
    422             }
    423 
    424             std::shared_ptr<IPackageVersion> packageVersion;
    425             if (!versionKey.Version.empty() && !versionKey.Channel.empty())
    426             {
    427                 for (const auto& versionInfo : m_package.Versions)
    428                 {
    429                     if (CaseInsensitiveEquals(versionInfo.VersionAndChannel.GetVersion().ToString(), versionKey.Version)
    430                         && CaseInsensitiveEquals(versionInfo.VersionAndChannel.GetChannel().ToString(), versionKey.Channel))
    431                     {
    432                         packageVersion = std::make_shared<PackageVersion>(source, NonConstSharedFromThis(), versionInfo);
    433                         break;
    434                     }
    435                 }
    436             }
    437             else if (versionKey.Version.empty() && versionKey.Channel.empty())
    438             {
    439                 packageVersion = GetLatestVersionInternal();
    440             }
    441             else if (versionKey.Version.empty())
    442             {
    443                 for (const auto& versionInfo : m_package.Versions)
    444                 {
    445                     if (CaseInsensitiveEquals(versionInfo.VersionAndChannel.GetChannel().ToString(), versionKey.Channel))
    446                     {
    447                         packageVersion = std::make_shared<PackageVersion>(source, NonConstSharedFromThis(), versionInfo);
    448                         break;
    449                     }
    450                 }
    451             }
    452             else if (versionKey.Channel.empty())
    453             {
    454                 for (const auto& versionInfo : m_package.Versions)
    455                 {
    456                     if (CaseInsensitiveEquals(versionInfo.VersionAndChannel.GetVersion().ToString(), versionKey.Version))
    457                     {
    458                         packageVersion = std::make_shared<PackageVersion>(source, NonConstSharedFromThis(), versionInfo);
    459                         break;
    460                     }
    461                 }
    462             }
    463 
    464             return packageVersion;
    465         }
    466 
    467         std::shared_ptr<IPackageVersion> RestPackage::GetLatestVersionInternal() const
    468         {
    469             return std::make_shared<PackageVersion>(GetReferenceSource(), NonConstSharedFromThis(), m_package.Versions.front());
    470         }
    471     }
    472 
    473     RestSource::RestSource(const SourceDetails& details, SourceInformation information, RestClient&& restClient)
    474         : m_details(details), m_information(std::move(information)), m_restClient(std::move(restClient))
    475     {
    476     }
    477 
    478     const std::string& RestSource::GetIdentifier() const
    479     {
    480         return m_details.Identifier;
    481     }
    482 
    483     const SourceDetails& RestSource::GetDetails() const
    484     {
    485         return m_details;
    486     }
    487 
    488     SourceInformation RestSource::GetInformation() const
    489     {
    490         return m_information;
    491     }
    492 
    493     bool RestSource::QueryFeatureFlag(SourceFeatureFlag flag) const
    494     {
    495         switch (flag)
    496         {
    497         case SourceFeatureFlag::ManifestMayContainAdditionalSystemReferenceStrings:
    498             return true;
    499         }
    500 
    501         return false;
    502     }
    503 
    504     SearchResult RestSource::Search(const SearchRequest& request) const
    505     {
    506         IRestClient::SearchResult results = m_restClient.Search(request);
    507         SearchResult searchResult;
    508 
    509         std::shared_ptr<RestSource> sharedThis = NonConstSharedFromThis();
    510         for (auto& result : results.Matches)
    511         {
    512             std::shared_ptr<ICompositePackage> package = std::make_shared<RestPackage>(sharedThis, std::move(result));
    513 
    514             // TODO: Improve to use Package match filter to return relevant search results.
    515             PackageMatchFilter packageFilter{ {}, {}, {} };
    516 
    517             searchResult.Matches.emplace_back(std::move(package), std::move(packageFilter));
    518         }
    519 
    520         searchResult.Truncated = results.Truncated;
    521 
    522         return searchResult;
    523     }
    524 
    525     void* RestSource::CastTo(ISourceType type)
    526     {
    527         if (type == SourceType)
    528         {
    529             return this;
    530         }
    531 
    532         return nullptr;
    533     }
    534 
    535     const RestClient& RestSource::GetRestClient() const
    536     {
    537         return m_restClient;
    538     }
    539 
    540     bool RestSource::IsSame(const RestSource* other) const
    541     {
    542         return (other && GetIdentifier() == other->GetIdentifier());
    543     }
    544 
    545     std::shared_ptr<RestSource> RestSource::NonConstSharedFromThis() const
    546     {
    547         return const_cast<RestSource*>(this)->shared_from_this();
    548     }
    549 }