diff --git a/CMakeLists.txt b/CMakeLists.txt index b074e6b0..4049607d 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -75,13 +75,18 @@ target_sources(${PROJECT_NAME} src/Model.h src/windows/AboutWindow.h - src/windows/WelcomeWindow.h + src/windows/CustomPathWindow.h + + src/windows/tutorial/TutorialTargets.h + src/windows/tutorial/TutorialOverlay.h + src/windows/tutorial/TutorialWindow.h + src/windows/settings/SettingsWindow.h src/windows/settings/GeneralSettingsTab.cpp src/windows/settings/LoginTab.cpp - src/widgets/ModelSelectionWidget.h src/widgets/ModelInfoWidget.h + src/widgets/ModelStyle.h src/widgets/ControlAreaWidget.h src/widgets/TrackAreaWidget.h src/widgets/StatusAreaWidget.h @@ -118,7 +123,8 @@ target_sources(${PROJECT_NAME} src/utils/Messages.h src/utils/Settings.h src/utils/Interface.h - src/utils/ModelRegistry.h + src/utils/ModelCatalog.h + src/utils/ModelTags.h src/utils/Controls.h src/utils/Labels.h src/utils/Clients.h @@ -185,6 +191,14 @@ juce_add_binary_data(stability_controls ${CMAKE_SOURCE_DIR}/src/clients/providers/stability/models/audio-to-audio.json ) +# The model taxonomy is defined once, in pyharp, where models declare their tags +juce_add_binary_data(model_taxonomy + HEADER_NAME TaxonomyData.h + NAMESPACE TaxonomyData + SOURCES + ${CMAKE_SOURCE_DIR}/pyharp/pyharp/taxonomy.json +) + # `target_link_libraries` links libraries and JUCE modules to other libraries or executables. Here, # we're linking our executable target to the `juce::juce_gui_extra` module. Inter-module # dependencies are resolved automatically, so `juce_core`, `juce_events` and so on will also be @@ -206,6 +220,7 @@ target_link_libraries(${PROJECT_NAME} juce::juce_gui_basics juce::juce_gui_extra stability_controls + model_taxonomy PUBLIC juce::juce_recommended_config_flags juce::juce_recommended_lto_flags diff --git a/README.md b/README.md index 74ad7058..373db679 100644 --- a/README.md +++ b/README.md @@ -77,8 +77,8 @@ HARP supports a simple workflow: pick an existing model for processing or provid To get started: - Open HARP as a standalone application or within your DAW -- Select an existing model using the drop-down menu at the top of the screen, or select `custom path...` and provide a URL to any HARP-compatible Gradio endpoint -- Load the selected model (and its corresponding interface) using the `Load` button. +- Browse or search the models on the `Home` tab, which lists every model hosted by [TEAMuP on Hugging Face](https://huggingface.co/teamup-tech) by category, or click `Custom Path...` and provide a URL to any HARP-compatible Gradio endpoint +- Click a model to open it (and its corresponding interface) in a tab of its own. Several models can be open at once, each in its own tab. - Import audio or MIDI data to process with the model either via the `Open File` button or by dragging and dropping a file into HARP - Adjust controls to taste in the interface - Click `Process` to run the model; outputs will automatically be rendered in HARP diff --git a/pyharp b/pyharp index 9660196e..77c0a78b 160000 --- a/pyharp +++ b/pyharp @@ -1 +1 @@ -Subproject commit 9660196ed711891712fb1ebf08a74d4b6cf1619b +Subproject commit 77c0a78b776dea267b97ad1426b459df34444b66 diff --git a/src/Application.cpp b/src/Application.cpp index 058e4fbc..d35ac15b 100644 --- a/src/Application.cpp +++ b/src/Application.cpp @@ -286,7 +286,7 @@ bool MainComponent::perform(const InvocationInfo& info) case CommandIDs::tutorial: DBG_AND_LOG("MainComponent::perform: \"tutorial\" command invoked."); - openWelcomeWindow(); + openTutorial(); break; diff --git a/src/HomeTab.h b/src/HomeTab.h index f70f921b..6fbfc1c0 100644 --- a/src/HomeTab.h +++ b/src/HomeTab.h @@ -1,6 +1,7 @@ /** * @file HomeTab.h - * @brief Home tab for model discovery and loading. + * @brief Home tab for browsing the model catalog and opening models in new tabs. + * @author JEYuhas, 2cylu2, VedMistry42, cwitkowitz */ #pragma once @@ -11,607 +12,758 @@ #include +#include "gui/HoverHandler.h" + #include "utils/Interface.h" -#include "utils/ModelRegistry.h" -#include "widgets/ModelSelectionWidget.h" +#include "utils/Messages.h" +#include "utils/ModelCatalog.h" +#include "utils/ModelTags.h" + +#include "widgets/ModelStyle.h" + +#include "windows/tutorial/TutorialTargets.h" + +#include "windows/CustomPathWindow.h" using namespace juce; -class TagLabel : public Component +/** + * Clickable card summarizing one catalog entry. + */ +class ModelCard : public Button { public: - TagLabel(const String& text) : tagText(text) + ModelCard(const CatalogEntry& catalogEntry, std::function onOpen) + : Button(catalogEntry.name), + entry(catalogEntry), + badges(ModelStyle::getBadges(catalogEntry)), + tagLabels(catalogEntry.tags.getDisplayLabels()), + searchableText(catalogEntry.getSearchableText()), + onOpenRequested(std::move(onOpen)) { - setSize(getPreferredWidth(), getPreferredHeight()); + setMouseCursor(MouseCursor::PointingHandCursor); } - int getPreferredWidth() const { return font.getStringWidth(tagText) + horizontalPadding * 2; } - int getPreferredHeight() const { return roundToInt(font.getHeight()) + verticalPadding * 2; } + const CatalogEntry& getEntry() const { return entry; } - void paint(Graphics& g) override + bool matchesSearch(const String& lowercaseQuery) const { - auto bounds = getLocalBounds().toFloat().reduced(0.5f); - g.setColour(Colour(0xff183238)); - g.fillRoundedRectangle(bounds, 4.0f); - g.setColour(Colour(0xff2dd4bf).withAlpha(0.42f)); - g.drawRoundedRectangle(bounds, 4.0f, 1.0f); - - g.setColour(Colour(0xff9eeadf)); - g.setFont(font); - g.drawText(tagText, getLocalBounds(), Justification::centred, true); + return lowercaseQuery.isEmpty() || searchableText.contains(lowercaseQuery); } -private: - String tagText; - Font font { 10.0f, Font::bold }; - static constexpr int horizontalPadding = 7; - static constexpr int verticalPadding = 4; -}; + // Removing a custom path is offered from the context menu + std::function onRemoveRequested; -class CategoryChip : public Button -{ -public: - CategoryChip(const String& name, bool selected) - : Button(name), isSelected(selected) + // Opens the model on a click, or offers to remove a custom path on a right-click + using Button::clicked; + void clicked(const ModifierKeys& modifiers) override { - } + if (! modifiers.isPopupMenu()) + { + if (onOpenRequested) + onOpenRequested(entry); - void setSelected(bool selected) - { - if (isSelected != selected) + return; + } + + if (entry.isCustom && onRemoveRequested) { - isSelected = selected; - repaint(); + PopupMenu menu; + menu.addItem("Remove from list", + [safeThis = SafePointer(this)] + { + if (safeThis != nullptr) + safeThis->onRemoveRequested(safeThis->entry); + }); + menu.showMenuAsync(PopupMenu::Options().withTargetComponent(this)); } } - bool getSelected() const { return isSelected; } + // Everything the card has no room for is shown in the instructions box while hovering + void mouseEnter(const MouseEvent& e) override + { + Button::mouseEnter(e); + + StringArray lines; + + if (entry.description.isNotEmpty()) + lines.add(entry.description); - int getPreferredWidth() const { return font.getStringWidth(getName()) + horizontalPadding * 2; } - int getPreferredHeight() const { return roundToInt(font.getHeight()) + verticalPadding * 2; } + if (const String tags = ModelStyle::describeTags(entry.tags); tags.isNotEmpty()) + lines.add(tags); - void paintButton(Graphics& g, bool shouldDrawButtonAsHighlighted, bool shouldDrawButtonAsDown) override + lines.add("Click to open " + entry.path + " in a new tab." + + (entry.isCustom ? " Right-click to remove it from the list." : "")); + + instructionsMessage->setMessage(lines.joinIntoString("\n")); + } + + void mouseExit(const MouseEvent& e) override { - auto bounds = getLocalBounds().toFloat().reduced(1.0f); - - Colour bg; - Colour textColour; - - if (isSelected) - { - bg = Colour(0xff0f766e); - textColour = Colours::white; - } - else if (shouldDrawButtonAsHighlighted || shouldDrawButtonAsDown) - { - bg = Colour(0xff263a3d); - textColour = Colours::white; - } - else - { - bg = Colour(0xff1e1e24); - textColour = Colours::lightgrey; - } + Button::mouseExit(e); + instructionsMessage->clearMessage(); + } - g.setColour(bg); - g.fillRoundedRectangle(bounds, 5.0f); + void paintButton(Graphics& g, bool isHighlighted, bool isDown) override + { + using namespace ModelStyle; + + drawCard(g, getLocalBounds().toFloat().reduced(1.0f), isHighlighted, isDown); + + auto area = getLocalBounds().reduced(12, 9); + + /* Name, with badges on the right */ + + auto nameRow = area.removeFromTop(20); + const int badgesLeft = drawBadges(g, badges, nameRow); + + g.setColour(Colours::white); + g.setFont(font(15.0f, true)); + g.drawText( + entry.name, nameRow.withRight(badgesLeft - chipGap), Justification::centredLeft, true); + + /* Description, cut short (all of it is shown while hovering) */ - g.setColour(isSelected ? Colour(0xff5eead4) : Colours::white.withAlpha(0.1f)); - g.drawRoundedRectangle(bounds, 5.0f, 1.0f); + area.removeFromTop(2); - g.setColour(textColour); - g.setFont(font); - g.drawText(getName(), getLocalBounds().reduced(horizontalPadding, 0), Justification::centred, true); + g.setColour(Colours::whitesmoke.withAlpha(0.85f)); + g.setFont(font(13.0f)); + drawTruncatedText(g, + entry.description.isNotEmpty() ? entry.description + : String("No description provided."), + area.removeFromTop(30), + 2); + + /* Tags, as many as fit (all of them are described while hovering) */ + + area.removeFromTop(4); + drawTagRow(g, tagLabels, area.removeFromTop(chipHeight)); + + /* Path */ + + area.removeFromTop(4); + + g.setColour(Colours::grey); + g.setFont(font(11.0f)); + g.drawText(entry.path, area.removeFromTop(13), Justification::centredLeft, true); } + static constexpr int preferredHeight = 110; + static constexpr int minimumWidth = 300; + private: - bool isSelected = false; - Font font { 13.0f, Font::bold }; - static constexpr int horizontalPadding = 12; - static constexpr int verticalPadding = 7; + const CatalogEntry entry; + const std::vector badges; + const StringArray tagLabels; + const String searchableText; + + std::function onOpenRequested; + + SharedResourcePointer instructionsMessage; }; -class CategoryFilterBar : public Component +/** + * The catalog as a responsive grid of cards, in one section per category. + */ +class ModelGrid : public Component { public: - CategoryFilterBar(std::function onCategorySelectedCallback) - : onCategorySelected(std::move(onCategorySelectedCallback)) - { - categories = { - "All", - "Generation", - "Performance Rendering and Synthesis", - "Effects", - "Enhancement", - "Production", - "Source Separation", - "Analysis", - "Custom" - }; + // Stands for models that declare no category + static inline const String otherSectionId { "other" }; + + std::function onOpenRequested; + std::function onRemoveRequested; + + void setEntries(const std::vector& entries) + { + sections.clear(); + removeAllChildren(); - for (int i = 0; i < categories.size(); ++i) + auto addSection = [&](const String& id, const String& title) { - auto chip = std::make_unique(categories[i], i == 0); - chip->onClick = [this, category = categories[i]] + auto section = std::make_unique
(); + section->id = id; + section->title = title; + + for (const auto& entry : entries) { - selectCategory(category); - }; - addAndMakeVisible(*chip); - chips.push_back(std::move(chip)); - } + const bool belongs = id == otherSectionId ? ! entry.tags.isCategorized() + : entry.tags.isInCategory(id); + + if (! belongs) + continue; + + auto* card = section->cards.add(new ModelCard(entry, onOpenRequested)); + card->onRemoveRequested = onRemoveRequested; + addChildComponent(card); + } + + // A model that belongs to several categories is listed in each of them + sections.push_back(std::move(section)); + }; + + for (const auto& category : Taxonomy::getCategories()) + addSection(category.id, category.displayName); + + addSection(otherSectionId, "Other"); } - void selectCategory(const String& category) + // Number of models in a section, or across all sections for an empty id + int countEntries(const String& sectionId) const { - for (auto& chip : chips) + StringArray paths; + + for (const auto& section : sections) { - chip->setSelected(chip->getName() == category); + if (sectionId.isNotEmpty() && section->id != sectionId) + continue; + + for (auto* card : section->cards) + paths.addIfNotAlreadyThere(card->getEntry().path); } - if (onCategorySelected) - onCategorySelected(category); + return paths.size(); } - void resized() override + /** + * Shows the cards in a section (or all sections, for an empty id) that match a + * search, and returns how many distinct models are shown. The grid has to be laid out + * again afterwards. + */ + int applyFilter(const String& sectionId, const String& searchText) { - auto area = getLocalBounds(); - int x = 0; - int y = 0; - int spacingX = 6; - int spacingY = 6; - int rowHeight = 0; + const String query = searchText.trim().toLowerCase(); - for (auto& chip : chips) + StringArray shownPaths; + + for (auto& section : sections) { - int chipWidth = chip->getPreferredWidth(); - int chipHeight = chip->getPreferredHeight(); - - if (x + chipWidth > area.getWidth() && x > 0) + const bool sectionShown = sectionId.isEmpty() || section->id == sectionId; + + for (auto* card : section->cards) { - x = 0; - y += rowHeight + spacingY; - rowHeight = 0; + const bool shown = sectionShown && card->matchesSearch(query); + card->setVisible(shown); + + if (shown) + shownPaths.addIfNotAlreadyThere(card->getEntry().path); } + } + + return shownPaths.size(); + } + + int getHeightForWidth(int width) const { return layOut(width, false); } + + void resized() override { layOut(getWidth(), true); } - chip->setBounds(x, y, chipWidth, chipHeight); - x += chipWidth + spacingX; - rowHeight = jmax(rowHeight, chipHeight); + void paint(Graphics& g) override + { + for (const auto& section : sections) + { + if (! section->headerBounds.isEmpty()) + ModelStyle::drawSectionHeader(g, section->title, section->headerBounds); } - - int newHeight = y + jmax(rowHeight, 1); - if (newHeight != preferredHeight) + } + +private: + struct Section + { + String id; + String title; + OwnedArray cards; + Rectangle headerBounds; + }; + + // Lays the visible cards out in as many columns as fit, and returns the height used + int layOut(int width, bool applyBounds) const + { + const int columns = jmax(1, (width + gap) / (ModelCard::minimumWidth + gap)); + const int cardWidth = jmax(0, (width - gap * (columns - 1)) / columns); + + int y = 0; + + for (const auto& section : sections) { - preferredHeight = newHeight; - MessageManager::callAsync([this]() + Array visibleCards; + + for (auto* card : section->cards) { - if (auto* parent = getParentComponent()) - parent->resized(); - }); + if (card->isVisible()) + visibleCards.add(card); + } + + if (visibleCards.isEmpty()) + { + if (applyBounds) + section->headerBounds = {}; + + continue; + } + + if (applyBounds) + section->headerBounds = { 0, y, width, ModelStyle::sectionHeaderHeight }; + + y += ModelStyle::sectionHeaderHeight; + + for (int i = 0; i < visibleCards.size(); ++i) + { + const int column = i % columns; + + if (column == 0 && i > 0) + y += ModelCard::preferredHeight + gap; + + if (applyBounds) + visibleCards[i]->setBounds( + column * (cardWidth + gap), y, cardWidth, ModelCard::preferredHeight); + } + + y += ModelCard::preferredHeight + sectionGap; } + + return y; } - int getPreferredHeight() const { return preferredHeight; } + static constexpr int gap = 8; + static constexpr int sectionGap = 10; -private: - std::vector categories; - std::vector> chips; - std::function onCategorySelected; - int preferredHeight = 28; + std::vector> sections; }; -class CategoryHeader : public Component +/** + * Chips for choosing which section of the catalog to show, wrapped onto as many rows as + * the width requires. + */ +class CategoryFilterBar : public Component { public: - CategoryHeader(const String& name) : categoryName(name) {} + std::function onSelectionChanged; - void paint(Graphics& g) override + CategoryFilterBar() { - auto bounds = getLocalBounds().toFloat(); - - g.setColour(Colours::white); - g.setFont(Font(16.0f, Font::bold)); - g.drawText(categoryName, getLocalBounds().reduced(4, 0), Justification::centredLeft, true); - - auto textWidth = Font(16.0f, Font::bold).getStringWidth(categoryName); - g.setColour(Colour(0xff2dd4bf).withAlpha(0.6f)); - g.fillRect(textWidth + 12.0f, bounds.getCentreY() - 1.0f, bounds.getWidth() - textWidth - 16.0f, 2.0f); - } + addChip({}, "All"); - static constexpr int preferredHeight = 32; + for (const auto& category : Taxonomy::getCategories()) + addChip(category.id, category.displayName); -private: - String categoryName; -}; + addChip(ModelGrid::otherSectionId, "Other"); -class ModelRegistryCard : public Component -{ -public: - ModelRegistryCard(ModelRegistry::Entry registryEntry, - std::function loadCallback) - : entry(std::move(registryEntry)), onLoad(std::move(loadCallback)) - { - nameLabel.setText(entry.displayName, dontSendNotification); - nameLabel.setJustificationType(Justification::centredLeft); - nameLabel.setFont(Font(17.0f, Font::bold)); - addAndMakeVisible(nameLabel); - - providerLabel.setText(entry.provider, dontSendNotification); - providerLabel.setJustificationType(Justification::centredLeft); - providerLabel.setColour(Label::textColourId, Colours::lightgrey); - addAndMakeVisible(providerLabel); - - summaryLabel.setText(entry.summary, dontSendNotification); - summaryLabel.setJustificationType(Justification::centredLeft); - summaryLabel.setColour(Label::textColourId, Colours::whitesmoke); - addAndMakeVisible(summaryLabel); - - pathLabel.setText(entry.path, dontSendNotification); - pathLabel.setJustificationType(Justification::centredLeft); - pathLabel.setColour(Label::textColourId, Colours::grey); - addAndMakeVisible(pathLabel); - - loadButton.setButtonText("Load"); - loadButton.onClick = [this] - { - if (onLoad) - onLoad(entry); - }; - addAndMakeVisible(loadButton); + chips.getFirst()->setToggleState(true, dontSendNotification); + } - for (const auto& tag : entry.tags) + // Id of the selected section, or empty for all of them + String getSelectedId() const + { + for (auto* chip : chips) { - auto label = std::make_unique(tag); - addAndMakeVisible(*label); - tagLabels.push_back(std::move(label)); + if (chip->getToggleState()) + return chip->id; } + + return {}; } - void paint(Graphics& g) override + void setCounts(const ModelGrid& grid) { - auto bounds = getLocalBounds().toFloat().reduced(1.0f); - g.setColour(getUIColourIfAvailable(LookAndFeel_V4::ColourScheme::UIColour::widgetBackground) - .brighter(0.06f)); - g.fillRoundedRectangle(bounds, 6.0f); + for (auto* chip : chips) + chip->count = grid.countEntries(chip->id); - g.setColour(Colours::white.withAlpha(0.12f)); - g.drawRoundedRectangle(bounds, 6.0f, 1.0f); + resized(); + repaint(); } - void resized() override + int getHeightForWidth(int width) const { - auto area = getLocalBounds().reduced(12, 10); - auto buttonArea = area.removeFromRight(92); - loadButton.setBounds(buttonArea.withSizeKeepingCentre(80, 30)); + FlexBox box = createLayout(); + box.performLayout(Rectangle(0, 0, width, 1000)); - auto topRow = area.removeFromTop(18); - providerLabel.setBounds(topRow.removeFromLeft(150)); - - for (auto& tagLabel : tagLabels) - { - tagLabel->setBounds(topRow.removeFromRight(tagLabel->getPreferredWidth() + 4) - .withSizeKeepingCentre(tagLabel->getPreferredWidth(), - tagLabel->getPreferredHeight())); - } + float bottom = 0.0f; + + for (const auto& item : box.items) + bottom = jmax(bottom, item.currentBounds.getBottom() + item.margin.bottom); - nameLabel.setBounds(area.removeFromTop(24)); - summaryLabel.setBounds(area.removeFromTop(24)); - pathLabel.setBounds(area.removeFromTop(18)); + return roundToInt(bottom); } - static constexpr int preferredHeight = 104; + void resized() override { createLayout().performLayout(getLocalBounds()); } private: - ModelRegistry::Entry entry; - std::function onLoad; - - Label nameLabel; - Label providerLabel; - Label summaryLabel; - Label pathLabel; - TextButton loadButton; - std::vector> tagLabels; -}; - -class ModelRegistryList : public Component -{ -public: - struct Section + struct Chip : public Button { - String category; - std::vector entries; - }; + Chip(const String& chipId, const String& chipName) : Button(chipName), id(chipId) + { + setClickingTogglesState(true); + setRadioGroupId(1); + } - void setSections(std::vector
newSections, - std::function loadCallback) - { - items.clear(); - removeAllChildren(); + String getText() const { return getName() + " " + String(count); } - for (auto& sec : newSections) + int getPreferredWidth() const { - if (sec.entries.empty()) - continue; + return ModelStyle::getTextWidth(chipFont, getText()) + 2 * horizontalPadding; + } - auto header = std::make_unique(sec.category); - addAndMakeVisible(*header); - items.push_back(std::move(header)); + void paintButton(Graphics& g, bool isHighlighted, bool isDown) override + { + const auto bounds = getLocalBounds().toFloat().reduced(1.0f); + const bool selected = getToggleState(); - for (auto& entry : sec.entries) - { - auto card = std::make_unique(std::move(entry), loadCallback); - addAndMakeVisible(*card); - items.push_back(std::move(card)); - } + g.setColour(selected ? ModelStyle::accentDark + : Colour(isHighlighted || isDown ? 0xff263a3d : 0xff1e1e24)); + g.fillRoundedRectangle(bounds, 5.0f); + + g.setColour(selected ? ModelStyle::accent : Colours::white.withAlpha(0.1f)); + g.drawRoundedRectangle(bounds, 5.0f, 1.0f); + + g.setColour(selected || isHighlighted ? Colours::white : Colours::lightgrey); + g.setFont(chipFont); + g.drawText(getText(), getLocalBounds(), Justification::centred, false); } - resized(); - repaint(); + void mouseEnter(const MouseEvent& e) override + { + Button::mouseEnter(e); + + if (id.isEmpty()) + instructionsMessage->setMessage("Click to show every model."); + else if (id == ModelGrid::otherSectionId) + instructionsMessage->setMessage( + "Click to show the models that declare no category."); + else + instructionsMessage->setMessage("Click to show only the " + getName() + " models."); + } + + void mouseExit(const MouseEvent& e) override + { + Button::mouseExit(e); + instructionsMessage->clearMessage(); + } + + const String id; + int count = 0; + + const Font chipFont = ModelStyle::font(12.0f, true); + static constexpr int horizontalPadding = 10; + + SharedResourcePointer instructionsMessage; + }; + + void addChip(const String& id, const String& name) + { + auto* chip = chips.add(new Chip(id, name)); + chip->onClick = [this] + { + if (onSelectionChanged) + onSelectionChanged(); + }; + addAndMakeVisible(chip); } - void resized() override + FlexBox createLayout() const { - auto area = getLocalBounds(); + FlexBox box; + box.flexWrap = FlexBox::Wrap::wrap; + box.alignContent = FlexBox::AlignContent::flexStart; - for (auto& item : items) + for (auto* chip : chips) { - if (dynamic_cast(item.get())) - item->setBounds(area.removeFromTop(CategoryHeader::preferredHeight)); - else if (dynamic_cast(item.get())) - item->setBounds(area.removeFromTop(ModelRegistryCard::preferredHeight).reduced(0, 4)); + box.items.add(FlexItem(*chip) + .withWidth((float) chip->getPreferredWidth()) + .withHeight((float) chipHeight) + .withMargin(FlexItem::Margin(0, 6, 6, 0))); } + + return box; + } + + static constexpr int chipHeight = 26; + + OwnedArray chips; +}; + +/** + * A small note on the state of the catalog (e.g., "Offline" or "2 hidden"), shown only when + * there is something to note, and explained in the instructions box while hovering over it. + */ +class CatalogStatusIndicator : public Component +{ +public: + void setStatus(const String& newText, Colour newColour, const String& newDetails) + { + text = newText; + colour = newColour; + details = newDetails; + + setVisible(text.isNotEmpty()); + repaint(); } - int getRequiredHeight() const + int getPreferredWidth() const { - int height = 0; - for (const auto& item : items) - { - if (dynamic_cast(item.get())) - height += CategoryHeader::preferredHeight; - else if (dynamic_cast(item.get())) - height += ModelRegistryCard::preferredHeight; - } - return height; + return text.isEmpty() ? 0 : ModelStyle::getTextWidth(textFont, text) + 16; + } + + void paint(Graphics& g) override + { + const auto bounds = getLocalBounds().toFloat().reduced(0.5f, 3.0f); + + g.setColour(colour.withAlpha(0.18f)); + g.fillRoundedRectangle(bounds, bounds.getHeight() / 2.0f); + + g.setColour(colour); + g.setFont(textFont); + g.drawText(text, getLocalBounds(), Justification::centred, false); } + void mouseEnter(const MouseEvent&) override { instructionsMessage->setMessage(details); } + + void mouseExit(const MouseEvent&) override { instructionsMessage->clearMessage(); } + private: - std::vector> items; + String text; + Colour colour; + String details; + + const Font textFont = ModelStyle::font(12.0f, true); + + SharedResourcePointer instructionsMessage; }; -class HomeTab : public Component, - private ChangeListener +class HomeTab : public Component, private ChangeListener { public: HomeTab() { - sharedChoices->addChangeListener(this); - titleLabel.setText("Models", dontSendNotification); - titleLabel.setJustificationType(Justification::centredLeft); - titleLabel.setFont(Font(24.0f, Font::bold)); + titleLabel.setFont(ModelStyle::font(20.0f, true)); + titleLabel.setBorderSize({ 0, 0, 0, 0 }); + addAndMakeVisible(titleLabel); + + addChildComponent(statusIndicator); - subtitleLabel.setText("Search HARP-compatible models and open one in a new tab.", - dontSendNotification); - subtitleLabel.setJustificationType(Justification::centredLeft); + addInstructions(refreshButton, + "Click to fetch the latest list of models from Hugging Face."); + refreshButton.onClick = [this] { catalog->refresh(); }; + addAndMakeVisible(refreshButton); - searchEditor.setTextToShowWhenEmpty("Search models...", Colours::grey); + addInstructions(customPathButton, + "Click to open a model that is not listed, by its Hugging Face Space, " + "Gradio URL, or local address."); + customPathButton.onClick = [this] + { + CustomPathComponent::launch(this, + [safeThis = SafePointer(this)](String path) + { + if (safeThis != nullptr) + safeThis->requestOpen(path, {}); + }); + }; + addAndMakeVisible(customPathButton); + + searchEditor.setTextToShowWhenEmpty("Search by name, task, tag, or path...", Colours::grey); searchEditor.setMultiLine(false); searchEditor.setReturnKeyStartsNewLine(false); - searchEditor.onTextChange = [this] { rebuildModelList(); }; + searchEditor.onTextChange = [this] { applyFilter(); }; + addInstructions(searchEditor, + "Type to show only the models whose name, description, tags, or path " + "contain the text."); + addAndMakeVisible(searchEditor); + + categoryFilterBar.onSelectionChanged = [this] { applyFilter(); }; + addAndMakeVisible(categoryFilterBar); - customPathButton.setButtonText("Custom Path"); - customPathButton.onClick = [this] { openCustomPathPopup(); }; + modelGrid.onOpenRequested = [this](const CatalogEntry& entry) + { requestOpen(entry.path, entry.name); }; + modelGrid.onRemoveRequested = [this](const CatalogEntry& entry) + { catalog->removeCustomPath(entry.path); }; - viewport.setViewedComponent(&modelList, false); + viewport.setViewedComponent(&modelGrid, false); viewport.setScrollBarsShown(true, false); - addAndMakeVisible(titleLabel); - addAndMakeVisible(subtitleLabel); - addAndMakeVisible(searchEditor); - addAndMakeVisible(customPathButton); - addAndMakeVisible(categoryFilterBar); + searchEditor.setComponentID(TutorialTargets::modelSearch); + viewport.setComponentID(TutorialTargets::modelList); addAndMakeVisible(viewport); - rebuildModelList(); + noResultsLabel.setJustificationType(Justification::centredTop); + noResultsLabel.setColour(Label::textColourId, Colours::grey); + addChildComponent(noResultsLabel); + + catalog->addChangeListener(this); + updateFromCatalog(); } ~HomeTab() override { - sharedChoices->removeChangeListener(this); + catalog->removeChangeListener(this); + + for (auto& handler : hoverHandlers) + handler->detach(); } + // Called with the path and display name of the model to open + std::function onModelOpenRequested; + void resized() override { - auto area = getLocalBounds().reduced(16); - - titleLabel.setBounds(area.removeFromTop(34)); - subtitleLabel.setBounds(area.removeFromTop(26)); + auto area = getLocalBounds().reduced(14, 12); + + FlexBox header; + header.alignItems = FlexBox::AlignItems::center; + header.items.add(FlexItem(titleLabel).withFlex(1).withHeight(28)); + header.items.add(FlexItem(statusIndicator) + .withWidth((float) statusIndicator.getPreferredWidth()) + .withHeight(26) + .withMargin({ 0, 8, 0, 0 })); + header.items.add( + FlexItem(refreshButton).withWidth(80).withHeight(26).withMargin({ 0, 6, 0, 0 })); + header.items.add(FlexItem(customPathButton).withWidth(110).withHeight(26)); + header.performLayout(area.removeFromTop(28)); area.removeFromTop(8); - auto searchRow = area.removeFromTop(34); - customPathButton.setBounds(searchRow.removeFromRight(120).reduced(0, 1)); - searchRow.removeFromRight(8); - searchEditor.setBounds(searchRow); + searchEditor.setBounds(area.removeFromTop(28)); - area.removeFromTop(10); - categoryFilterBar.setBounds(area.removeFromTop(categoryFilterBar.getPreferredHeight())); + area.removeFromTop(8); + categoryFilterBar.setBounds( + area.removeFromTop(categoryFilterBar.getHeightForWidth(area.getWidth()))); - area.removeFromTop(10); viewport.setBounds(area); + noResultsLabel.setBounds(area.withTrimmedTop(24)); - updateListBounds(); + layOutGrid(); } - void resetSelection() - { - searchEditor.setEnabled(true); - customPathButton.setEnabled(true); - categoryFilterBar.setEnabled(true); - viewport.setEnabled(true); - } +private: + void changeListenerCallback(ChangeBroadcaster*) override { updateFromCatalog(); } - Rectangle getModelSelectBounds() const + void updateFromCatalog() { - return searchEditor.getBounds().expanded(2, 2); - } + modelGrid.setEntries(catalog->getEntries()); + categoryFilterBar.setCounts(modelGrid); - std::function onModelLoadRequested; + updateStatusIndicator(); + refreshButton.setEnabled(catalog->getFetchState() != ModelCatalog::FetchState::Fetching); -private: - void changeListenerCallback(ChangeBroadcaster* source) override - { - if (source == static_cast(sharedChoices)) - rebuildModelList(); + // The chip counts, and so the rows they wrap onto, may have changed as well + updateVisibleModels(); + resized(); } - void requestModelLoad(const ModelRegistry::Entry& entry) + void updateStatusIndicator() { - searchEditor.setEnabled(false); - customPathButton.setEnabled(false); - categoryFilterBar.setEnabled(false); - viewport.setEnabled(false); + const String source = "huggingface.co/" + String(ModelCatalog::hubOrganization); + const auto& hidden = catalog->getHiddenEntries(); - if (onModelLoadRequested) - onModelLoadRequested(entry.path, entry.displayName); - } - - void rebuildModelList() - { - std::vector entries; - const auto searchText = searchEditor.getText().trim().toLowerCase(); + StringArray notes; + StringArray details; + Colour colour = ModelStyle::secondaryText; - for (const auto& savedPath : sharedChoices->savedModelPaths) + switch (catalog->getFetchState()) { - const String path(savedPath); - - if (path.startsWithIgnoreCase("click here")) - continue; - - auto entry = ModelRegistry::getEntryForPath(path); - const auto searchableText = - (entry.displayName + " " + entry.summary + " " + entry.path + " " + entry.provider) - .toLowerCase(); + case ModelCatalog::FetchState::Fetching: + notes.add("Updating..."); + details.add("Fetching the latest list of models from " + source + "."); + colour = ModelStyle::information; + break; - if (searchText.isEmpty() || searchableText.contains(searchText)) + case ModelCatalog::FetchState::Failed: { - if (activeCategory == "All") - { - entries.push_back(std::move(entry)); - } - else if (activeCategory == "Custom") - { - if (entry.tags.empty()) - entries.push_back(std::move(entry)); - } - else - { - bool matchesCategory = false; - for (const auto& tag : entry.tags) - { - if (tag == activeCategory) - { - matchesCategory = true; - break; - } - } - if (matchesCategory) - entries.push_back(std::move(entry)); - } + const Time listingTime = catalog->getListingTime(); + + notes.add("Offline"); + details.add("Could not reach " + source + " (" + catalog->getFetchError() + "). " + + (listingTime != Time() + ? "Showing the list from " + + listingTime.toString(true, true, false) + "." + : String("Only built-in and custom models are shown.")) + + " Click Refresh to try again."); + colour = ModelStyle::problem; + break; } - } - std::vector sections; - std::vector categoriesToShow; - - if (activeCategory == "All") - { - categoriesToShow = { - "Generation", - "Performance Rendering and Synthesis", - "Effects", - "Enhancement", - "Production", - "Source Separation", - "Analysis", - "Custom" - }; - } - else - { - categoriesToShow = { activeCategory }; + case ModelCatalog::FetchState::Succeeded: + break; } - for (const auto& cat : categoriesToShow) + if (! hidden.empty()) { - ModelRegistryList::Section sec; - sec.category = cat; - - for (const auto& entry : entries) - { - if (cat == "Custom") - { - if (entry.tags.empty()) - sec.entries.push_back(entry); - } - else - { - for (const auto& tag : entry.tags) - { - if (tag == cat) - { - sec.entries.push_back(entry); - break; - } - } - } - } - - if (! sec.entries.empty()) - sections.push_back(std::move(sec)); + notes.add(String((int) hidden.size()) + " hidden"); + + StringArray hiddenLines { "Hidden, since they cannot currently be loaded:" }; + + for (const auto& entry : hidden) + hiddenLines.add(entry.name + " (" + entry.reason + ")"); + + details.add(hiddenLines.joinIntoString("\n")); } - modelList.setSections(std::move(sections), - [this](ModelRegistry::Entry entry) { requestModelLoad(entry); }); - updateListBounds(); + statusIndicator.setStatus(notes.joinIntoString(String::fromUTF8(" \xc2\xb7 ")), + colour, + details.joinIntoString("\n")); } - void updateListBounds() + void applyFilter() { - const auto width = jmax(0, viewport.getWidth() - viewport.getScrollBarThickness()); - modelList.setSize(width, jmax(viewport.getHeight(), modelList.getRequiredHeight())); + updateVisibleModels(); + layOutGrid(); } - void openCustomPathPopup() + void updateVisibleModels() { - std::function loadCallback = [this](String path) - { - auto entry = ModelRegistry::getEntryForPath(path); - requestModelLoad(entry); - }; + const int numShown = + modelGrid.applyFilter(categoryFilterBar.getSelectedId(), searchEditor.getText()); + + noResultsLabel.setText(searchEditor.isEmpty() + ? "No models in this category yet." + : "No models match \"" + searchEditor.getText().trim() + "\".", + dontSendNotification); + noResultsLabel.setVisible(numShown == 0); + } - auto* content = new CustomPathComponent(std::move(loadCallback), [] {}); + void layOutGrid() + { + const int width = jmax(0, viewport.getWidth() - viewport.getScrollBarThickness()); + const int height = jmax(viewport.getHeight(), modelGrid.getHeightForWidth(width)); - DialogWindow::LaunchOptions options; - options.dialogTitle = "Enter Custom Path"; - options.dialogBackgroundColour = Colours::darkgrey; - options.content.setOwned(content); + // Which cards are shown can change without the size changing + if (modelGrid.getWidth() == width && modelGrid.getHeight() == height) + modelGrid.resized(); + else + modelGrid.setSize(width, height); - options.useNativeTitleBar = false; - options.resizable = false; - options.escapeKeyTriggersCloseButton = true; - options.componentToCentreAround = this; + modelGrid.repaint(); + } - options.launchAsync(); + void requestOpen(const String& path, const String& name) + { + if (onModelOpenRequested) + onModelOpenRequested(path, name); + } + + // Shows instructions for a component in the instructions box while hovering over it + void addInstructions(Component& component, const String& instructions) + { + auto handler = std::make_unique(component); + + handler->onMouseEnter = [this, instructions] + { instructionsMessage->setMessage(instructions); }; + handler->onMouseExit = [this] { instructionsMessage->clearMessage(); }; + handler->attach(); + + hoverHandlers.push_back(std::move(handler)); } Label titleLabel; - Label subtitleLabel; + CatalogStatusIndicator statusIndicator; + TextButton refreshButton { "Refresh" }; + TextButton customPathButton { "Custom Path..." }; TextEditor searchEditor; - TextButton customPathButton; - CategoryFilterBar categoryFilterBar { [this](String cat) { activeCategory = cat; rebuildModelList(); } }; - String activeCategory { "All" }; + CategoryFilterBar categoryFilterBar; Viewport viewport; - ModelRegistryList modelList; + ModelGrid modelGrid; + Label noResultsLabel; + + std::vector> hoverHandlers; - SharedResourcePointer sharedChoices; + SharedResourcePointer catalog; + SharedResourcePointer instructionsMessage; }; diff --git a/src/Main.cpp b/src/Main.cpp index 7ecf8618..6baac629 100644 --- a/src/Main.cpp +++ b/src/Main.cpp @@ -6,8 +6,6 @@ #include "MainComponent.h" -#include "windows/WelcomeWindow.h" - #include "utils/Settings.h" #if JUCE_LINUX @@ -87,7 +85,7 @@ class GuiAppApplication : public JUCEApplication, public FocusChangeListener if (auto* mainComp = dynamic_cast(getMainWindowPtr()->getContentComponent())) { - mainComp->openWelcomeWindow(); + mainComp->openTutorial(); } } diff --git a/src/MainComponent.cpp b/src/MainComponent.cpp index d9bb695a..bbfb7d13 100644 --- a/src/MainComponent.cpp +++ b/src/MainComponent.cpp @@ -1,7 +1,5 @@ #include "MainComponent.h" -#include "windows/WelcomeWindow.h" - JUCE_IMPLEMENT_SINGLETON(HARPLogger) MainComponent::MainComponent() @@ -12,10 +10,12 @@ MainComponent::MainComponent() modelTabs.addChangeListener(this); - modelTabs.onPageScrolled = [this] { refreshTutorialHighlight(); }; addAndMakeVisible(modelTabs); addAndMakeVisible(statusAreaWidget); addAndMakeVisible(mediaClipboardWidget); + + statusAreaWidget.setComponentID(TutorialTargets::statusArea); + mediaClipboardWidget.setComponentID(TutorialTargets::clipboard); addAndMakeVisible(dragOverlay); showStatusArea = Settings::getBoolValue("view.showStatusArea", true); @@ -43,56 +43,6 @@ void MainComponent::paint(Graphics& g) g.fillAll(getUIColourIfAvailable(LookAndFeel_V4::ColourScheme::UIColour::windowBackground)); } -void MainComponent::paintOverChildren(Graphics& g) -{ - if (isTutorialActive) - { - auto area = getLocalBounds(); - g.setColour(Colours::black.withAlpha(0.6f)); - - if (tutorialHighlightRect.isEmpty() && tutorialExtraHighlights.empty()) - { - // Full dim if no highlight - g.fillAll(); - } - else - { - // Dim with cutout - Path backgroundPath; - backgroundPath.addRectangle(area.toFloat()); - - Path highlightPath; - if (! tutorialHighlightRect.isEmpty()) - highlightPath.addRoundedRectangle(tutorialHighlightRect.toFloat(), 5.0f); - - // Add extra highlights to the cutout path - for (auto& rect : tutorialExtraHighlights) - { - if (! rect.isEmpty()) - highlightPath.addRoundedRectangle(rect.toFloat(), 5.0f); - } - - backgroundPath.setUsingNonZeroWinding(false); - backgroundPath.addPath(highlightPath); - - g.fillPath(backgroundPath); - - g.setColour(Colours::white); - - // An empty rectangle is a region the current step has nothing to - // point at; outlining it would leave a stray mark in the corner - if (! tutorialHighlightRect.isEmpty()) - g.drawRoundedRectangle(tutorialHighlightRect.toFloat(), 5.0f, 2.0f); - - for (auto& rect : tutorialExtraHighlights) - { - if (! rect.isEmpty()) - g.drawRoundedRectangle(rect.toFloat(), 5.0f, 2.0f); - } - } - } -} - void MainComponent::resized() { Rectangle fullArea = getLocalBounds(); @@ -132,36 +82,13 @@ void MainComponent::resized() fullWindow.performLayout(fullArea); - /* Deferred: the highlight is measured from component bounds, which are only - final once this layout pass and the tab's own have completed. */ - refreshTutorialHighlight(); - dragOverlay.setBounds(getLocalBounds()); } -void MainComponent::refreshTutorialHighlight() -{ - if (welcomeWindow == nullptr) - { - return; - } - - Component::SafePointer safeThis(this); - - MessageManager::callAsync( - [safeThis] - { - if (safeThis != nullptr && safeThis->welcomeWindow != nullptr) - { - safeThis->welcomeWindow->refreshHighlightForCurrentStep(); - } - }); -} - void MainComponent::updateWindowConstraints() { // The Home tab has no controls, so only the general minimums apply while it is showing - auto* tab = getCurrentModelTab(); + auto* tab = modelTabs.getCurrentModelTab(); const int requiredControlWidth = tab != nullptr ? tab->getMinimumRequiredControlWidth() : 0; if (auto* window = findParentComponentOfClass()) @@ -415,14 +342,11 @@ void MainComponent::openAboutWindow() options.launchAsync(); } -void MainComponent::openWelcomeWindow(bool ensureDefaultModelLoaded) +void MainComponent::openTutorial() { - if (ensureDefaultModelLoaded) - ensureTutorialModelLoaded(); - - if (welcomeWindow != nullptr) + if (tutorialWindow != nullptr) { - welcomeWindow->toFront(true); + tutorialWindow->toFront(true); return; } @@ -433,270 +357,26 @@ void MainComponent::openWelcomeWindow(bool ensureDefaultModelLoaded) if (safeThis == nullptr) return; - safeThis->welcomeWindow.reset(new WelcomeWindow(safeThis.getComponent())); - safeThis->welcomeWindow->onClose = [safeThis]() + TutorialHost& host = *safeThis.getComponent(); + safeThis->tutorialWindow = std::make_unique(host); + safeThis->tutorialWindow->onClose = [safeThis]() { if (safeThis != nullptr) - safeThis->welcomeWindow.reset(); + safeThis->tutorialWindow.reset(); }; - safeThis->welcomeWindow->setVisible(true); - safeThis->welcomeWindow->positionOnMainComponentDisplay(); - safeThis->welcomeWindow->toFront(true); - }); -} - -/* --Tutorial-- */ - -void MainComponent::setTutorialActive(bool active) -{ - isTutorialActive = active; - repaint(); -} - -void MainComponent::setTutorialHighlight(Rectangle bounds) -{ - tutorialHighlightRect = bounds; - repaint(); -} - -void MainComponent::setTutorialExtraHighlights(std::vector> bounds) -{ - tutorialExtraHighlights = bounds; - repaint(); -} - -void MainComponent::ensureTutorialModelLoaded() -{ - // Loading is asynchronous, so without this guard every repeated call that - // arrives before the first load finishes - clicking Next again, say - would - // open yet another tab or start yet another load. - if (tutorialModelLoadInFlight) - return; - - auto* tab = getCurrentModelTab(); - - if (tab == nullptr) - { - // createNewTab() selects the tab it creates, which is what the tutorial - // steps compute their highlights against; leave it selected. - tab = modelTabs.createNewTab(); - tutorialCreatedTab = tab; - } - - if (tab->isModelLoaded()) - return; - - tutorialModelLoadInFlight = true; - - Component::SafePointer safeThis(this); - tab->onNextModelLoadComplete( - [safeThis](ModelTab*, bool) - { - if (safeThis != nullptr) - safeThis->tutorialModelLoadInFlight = false; + safeThis->tutorialWindow->setVisible(true); + safeThis->tutorialWindow->positionOnHostDisplay(); + safeThis->tutorialWindow->toFront(true); }); - - tab->loadDefaultModel(); -} - -void MainComponent::resetTutorialAutoLoadedModel() -{ - // Close the tab the tutorial opened on the user's behalf. Resetting it in - // place would leave a blank tab behind, since a model tab has no model - // selection of its own - models are chosen on the Home tab. - if (auto* tab = tutorialCreatedTab.getComponent()) - modelTabs.closeTab(tab); - - tutorialCreatedTab = nullptr; } -void MainComponent::ensureMediaClipboardVisible() +void MainComponent::openMediaClipboard() { if (! showMediaClipboard) viewMediaClipboardCallback(); } -Rectangle MainComponent::getTabBarBounds() -{ - auto& tabBar = modelTabs.getTabbedButtonBar(); - - if (tabBar.getNumTabs() == 0) - return {}; - - return getLocalArea(&tabBar, tabBar.getLocalBounds()); -} - -/** - * Converts a rectangle from the current model tab's coordinates into this component's, - * clipped to the part of the tab its page is showing. - * - * The tab can be taller than its page, so a component that is scrolled out of view - * could otherwise produce a highlight lying over the status area beneath it. - */ -Rectangle MainComponent::getVisibleTabArea(Rectangle tabBounds) -{ - auto* page = modelTabs.getCurrentModelTabPage(); - - if (page == nullptr || tabBounds.isEmpty()) - return {}; - - return getLocalArea(&page->getModelTab(), tabBounds) - .getIntersection(getLocalArea(page, page->getLocalBounds())); -} - -Rectangle MainComponent::getModelSelectBounds() -{ - if (auto* homeTab = modelTabs.getHomeTabIfShowing()) - return getLocalArea(homeTab, homeTab->getModelSelectBounds()); - - // Models are selected on the Home tab, so while a model tab is showing, - // point at the tab bar that leads back to it. - return getTabBarBounds(); -} - -Rectangle MainComponent::getControlsBounds() -{ - if (auto* tab = getCurrentModelTab()) - return getVisibleTabArea(tab->getControlsBounds()); - - return {}; -} - -Rectangle MainComponent::getInputTrackBounds() -{ - if (auto* tab = getCurrentModelTab()) - return getVisibleTabArea(tab->getInputTrackBounds()); - - return {}; -} - -Rectangle MainComponent::getInputFolderBounds() -{ - if (auto* tab = getCurrentModelTab()) - return getVisibleTabArea(tab->getInputFolderBounds()); - - return {}; -} - -Rectangle MainComponent::getInputPlayBounds() -{ - if (auto* tab = getCurrentModelTab()) - return getVisibleTabArea(tab->getInputPlayBounds()); - - return {}; -} - -Rectangle MainComponent::getProcessButtonBounds() -{ - if (auto* tab = getCurrentModelTab()) - return getVisibleTabArea(tab->getProcessButtonBounds()); - - return {}; -} - -Rectangle MainComponent::getTracksBounds() -{ - if (auto* tab = getCurrentModelTab()) - return getVisibleTabArea(tab->getTracksBounds()); - - return {}; -} - -Rectangle MainComponent::getClipboardBounds() -{ - if (showMediaClipboard && mediaClipboardWidget.isVisible()) - return mediaClipboardWidget.getBounds(); - return {}; -} - -Rectangle MainComponent::getClipboardTrackAreaBounds() -{ - if (showMediaClipboard && mediaClipboardWidget.isVisible()) - { - auto bounds = mediaClipboardWidget.getClipboardTrackAreaBounds(); - return getLocalArea(&mediaClipboardWidget, bounds); - } - return {}; -} - -Rectangle MainComponent::getClipboardControlsBounds() -{ - if (showMediaClipboard && mediaClipboardWidget.isVisible()) - { - auto bounds = mediaClipboardWidget.getClipboardControlsBounds(); - return getLocalArea(&mediaClipboardWidget, bounds); - } - return {}; -} - -Rectangle MainComponent::getClipboardNameBoxBounds() -{ - if (showMediaClipboard && mediaClipboardWidget.isVisible()) - { - auto bounds = mediaClipboardWidget.getClipboardNameBoxBounds(); - return getLocalArea(&mediaClipboardWidget, bounds); - } - return {}; -} - -Rectangle MainComponent::getClipboardButtonsBounds() -{ - if (showMediaClipboard && mediaClipboardWidget.isVisible()) - { - auto bounds = mediaClipboardWidget.getClipboardButtonsBounds(); - return getLocalArea(&mediaClipboardWidget, bounds); - } - return {}; -} - -Rectangle MainComponent::getClipboardAddButtonBounds() -{ - if (showMediaClipboard && mediaClipboardWidget.isVisible()) - { - auto bounds = mediaClipboardWidget.getAddFileButtonBounds(); - return getLocalArea(&mediaClipboardWidget, bounds); - } - return {}; -} - -Rectangle MainComponent::getClipboardRemoveButtonBounds() -{ - if (showMediaClipboard && mediaClipboardWidget.isVisible()) - { - auto bounds = mediaClipboardWidget.getRemoveButtonBounds(); - return getLocalArea(&mediaClipboardWidget, bounds); - } - return {}; -} - -Rectangle MainComponent::getClipboardPlayButtonBounds() -{ - if (showMediaClipboard && mediaClipboardWidget.isVisible()) - { - auto bounds = mediaClipboardWidget.getPlayButtonBounds(); - return getLocalArea(&mediaClipboardWidget, bounds); - } - return {}; -} - -Rectangle MainComponent::getClipboardSendToDAWButtonBounds() -{ - if (showMediaClipboard && mediaClipboardWidget.isVisible()) - { - auto bounds = mediaClipboardWidget.getSendToDAWButtonBounds(); - return getLocalArea(&mediaClipboardWidget, bounds); - } - return {}; -} - -Rectangle MainComponent::getInfoBarBounds() -{ - if (showStatusArea && statusAreaWidget.isVisible()) - return statusAreaWidget.getBounds(); - return {}; -} - /* --Miscellaneous-- */ // TODO - The following is an old callback from V2. It may be helpful in the future. @@ -749,10 +429,5 @@ void MainComponent::changeListenerCallback(ChangeBroadcaster* source) if (source == &modelTabs) { updateWindowConstraints(); - - // Model tabs are created and closed while the tutorial is open, so it - // follows the container rather than subscribing to individual tabs. - if (welcomeWindow != nullptr) - welcomeWindow->notifyModelStateChanged(); } } diff --git a/src/MainComponent.h b/src/MainComponent.h index 6cb016c5..1f36071a 100644 --- a/src/MainComponent.h +++ b/src/MainComponent.h @@ -16,24 +16,22 @@ #include "widgets/MediaClipboardWidget.h" #include "widgets/StatusAreaWidget.h" +#include "windows/tutorial/TutorialWindow.h" + #include "windows/AboutWindow.h" #include "windows/settings/SettingsWindow.h" #include "utils/Interface.h" #include "utils/Logging.h" #include "utils/Settings.h" -#include "utils/Tutorial.h" using namespace juce; -// Forward declaration (include in .cpp) -class WelcomeWindow; - class MainComponent : public Component, public MenuBarModel, public ApplicationCommandTarget, - private ChangeListener - + private ChangeListener, + private TutorialHost { public: MainComponent(); @@ -72,49 +70,14 @@ class MainComponent : public Component, // Help void openAboutWindow(); - void openWelcomeWindow(bool ensureDefaultModelLoaded = false); - - /* Tutorial */ - - ModelTab* getCurrentModelTab() const { return modelTabs.getCurrentModelTab(); } - - void setTutorialActive(bool active); - void setTutorialHighlight(Rectangle bounds); - void setTutorialExtraHighlights(std::vector> bounds); - void ensureTutorialModelLoaded(); - bool isTutorialModelLoadInFlight() const { return tutorialModelLoadInFlight; } - void resetTutorialAutoLoadedModel(); - void ensureMediaClipboardVisible(); - - // Bounds accessors for tutorial steps (public for WelcomeWindow) - Rectangle getTabBarBounds(); - Rectangle getModelSelectBounds(); - Rectangle getControlsBounds(); - Rectangle getInputTrackBounds(); - Rectangle getInputFolderBounds(); - Rectangle getInputPlayBounds(); - Rectangle getProcessButtonBounds(); - Rectangle getTracksBounds(); - Rectangle getClipboardBounds(); - Rectangle getClipboardTrackAreaBounds(); - Rectangle getClipboardControlsBounds(); - Rectangle getClipboardNameBoxBounds(); - Rectangle getClipboardButtonsBounds(); - Rectangle getClipboardAddButtonBounds(); - Rectangle getClipboardRemoveButtonBounds(); - Rectangle getClipboardPlayButtonBounds(); - Rectangle getClipboardSendToDAWButtonBounds(); - Rectangle getInfoBarBounds(); + void openTutorial(); /* Component */ void paint(Graphics& g) override; - void paintOverChildren(Graphics& g) override; void resized() override; void updateWindowConstraints(); - void refreshTutorialHighlight(); - Rectangle getVisibleTabArea(Rectangle tabBounds); private: /* File Menu */ @@ -135,6 +98,12 @@ class MainComponent : public Component, //void focusCallback(); void changeListenerCallback(ChangeBroadcaster* source) override; + /* Tutorial */ + + Component& getTutorialArea() override { return *this; } + ModelTabContainer& getModelTabs() override { return modelTabs; } + void openMediaClipboard() override; + /* Interface */ const int statusAreaHeight = 100; @@ -142,7 +111,7 @@ class MainComponent : public Component, const float mediaClipboardScale = 1.4f; // Minimum size to ensure all controls remain visible and functional: - // - WelcomeWindow popup is 480x500, needs padding + // - Tutorial window is 500x420, needs padding // - Dropdown labels need adequate width // - Control Area needs space for sliders/toggles/textboxes const int minimumWindowWidth = 700; @@ -164,16 +133,7 @@ class MainComponent : public Component, DragOverlayComponent dragOverlay; MediaClipboardWidget mediaClipboardWidget { &dragOverlay }; - bool isTutorialActive = false; - Rectangle tutorialHighlightRect; - std::vector> tutorialExtraHighlights; - std::unique_ptr welcomeWindow; - - // Set while the tutorial's fallback model is loading, so that repeated - // requests to load it do not stack up. The tab the tutorial opened on the - // user's behalf is remembered so that it can be closed again at the end. - bool tutorialModelLoadInFlight = false; - Component::SafePointer tutorialCreatedTab; + std::unique_ptr tutorialWindow; SharedResourcePointer sharedTokens; SharedResourcePointer statusMessage; diff --git a/src/ModelTab.h b/src/ModelTab.h index 74e4e83b..a0a9d638 100644 --- a/src/ModelTab.h +++ b/src/ModelTab.h @@ -7,7 +7,6 @@ #pragma once #include -#include #include @@ -15,115 +14,76 @@ #include "widgets/ControlAreaWidget.h" #include "widgets/ModelInfoWidget.h" -#include "widgets/ModelSelectionWidget.h" #include "widgets/TrackAreaWidget.h" #include "utils/Errors.h" #include "utils/Logging.h" +#include "utils/ModelCatalog.h" #include "utils/Settings.h" -#include "utils/Tutorial.h" + +#include "windows/tutorial/TutorialTargets.h" using namespace juce; -class ModelTab : public Component, private ChangeListener, public ChangeBroadcaster +/** + * Sends a change message once a model has loaded, and once the error from a failed load + * has been dismissed. A tab whose first load failed is then left without a model. + */ +class ModelTab : public Component, public ChangeBroadcaster { public: ModelTab() { - modelSelectionWidget.addChangeListener(this); - addAndMakeVisible(modelInfoWidget); addAndMakeVisible(controlAreaWidget); - inputTracksLabel.setJustificationType(Justification::centred); - inputTracksLabel.setFont(Font(20.0f, Font::bold)); - addAndMakeVisible(inputTracksLabel); addAndMakeVisible(inputTrackAreaWidget); initializeProcessCancelButton(); - outputTracksLabel.setJustificationType(Justification::centred); - outputTracksLabel.setFont(Font(20.0f, Font::bold)); - addAndMakeVisible(outputTracksLabel); addAndMakeVisible(outputTrackAreaWidget); - } - ~ModelTab() { modelSelectionWidget.removeChangeListener(this); } - - // Accessor methods for WelcomeWindow tutorial - std::shared_ptr getModel() const { return model; } - String getLoadedPath() const { return model->getLoadedPath(); } - - void loadDefaultModel() - { - modelSelectionWidget.loadModelBypass(TutorialConstants::fallbackModelPath); - } - - void loadModelPath(const String& modelPath) - { - modelSelectionWidget.loadModelBypass(modelPath); + controlAreaWidget.setComponentID(TutorialTargets::controls); + inputTrackAreaWidget.setComponentID(TutorialTargets::inputTracks); + outputTrackAreaWidget.setComponentID(TutorialTargets::outputTracks); + processCancelButton.setComponentID(TutorialTargets::processButton); } - void onNextModelLoadComplete(std::function callback) - { - initialLoadCallback = std::move(callback); - } + std::shared_ptr getModel() const { return model; } - // Bounds accessors for tutorial steps - Rectangle getModelSelectBounds() const + // Loads a model in the background, reporting the outcome with a change message + void loadModel(const String& modelPath) { - auto bounds = modelSelectionWidget.getBounds(); - - // The model browser lives on the Home tab, so this widget is laid out - // with an empty size here. Report nothing rather than a stray rectangle - // in the top left corner. - if (bounds.getWidth() > 0 && bounds.getHeight() > 0) - return bounds.expanded(2, 2); + loading = true; - return {}; - } + // Disable processing until model is loaded + processCancelButton.setEnabled(false); - Rectangle getControlsBounds() const - { - auto bounds = controlAreaWidget.getBounds(); + const String pathToLoad = canonicalizeModelPath(modelPath); - if (bounds.getWidth() > 0 && bounds.getHeight() > 0) - return bounds.expanded(2, 2); + DBG_AND_LOG("ModelTab::loadModel: Attempting to load path \"" << pathToLoad << "\"."); - return {}; - } + SafePointer safeThis(this); - Rectangle getInputFolderBounds() - { - auto bounds = inputTrackAreaWidget.getFirstTrackFolderButtonBounds(); - return getLocalArea(&inputTrackAreaWidget, bounds); - } + loadingThreadPool.addJob( + [this, safeThis, pathToLoad] + { + OpResult result = model->load(pathToLoad); - Rectangle getInputPlayBounds() - { - auto bounds = inputTrackAreaWidget.getFirstTrackPlayButtonBounds(); - return getLocalArea(&inputTrackAreaWidget, bounds); + // Perform updates on message (GUI) thread + MessageManager::callAsync( + [safeThis, result, pathToLoad] + { + if (safeThis != nullptr) + safeThis->finishLoading(result, pathToLoad); + }); + }); } - Rectangle getInputTrackBounds() const { return inputTrackAreaWidget.getBounds(); } - - Rectangle getProcessButtonBounds() const { return processCancelButton.getBounds(); } - - Rectangle getTracksBounds() const - { - auto bounds = inputTrackAreaWidget.getBounds(); - if (outputTrackAreaWidget.isVisible()) - bounds = bounds.getUnion(outputTrackAreaWidget.getBounds()); - - if (inputTracksLabel.isVisible()) - bounds = bounds.getUnion(inputTracksLabel.getBounds()); - if (outputTracksLabel.isVisible()) - bounds = bounds.getUnion(outputTracksLabel.getBounds()); - - return bounds.expanded(2, 2); - } + // True from the moment a load is requested until its outcome has been reported + bool isLoading() const { return loading; } bool isModelLoaded() { return model->isLoaded(); } @@ -158,11 +118,8 @@ class ModelTab : public Component, private ChangeListener, public ChangeBroadcas FlexBox tabArea; tabArea.flexDirection = FlexBox::Direction::column; - const int width = getWidth(); - - /* Model Selection */ - - modelSelectionWidget.setBounds(0, 0, 0, 0); + const auto contentArea = getLocalBounds().reduced(pagePadding); + const int width = contentArea.getWidth(); /* Model Info */ @@ -232,7 +189,7 @@ class ModelTab : public Component, private ChangeListener, public ChangeBroadcas outputTrackAreaWidget.getNumTracks(), totalTracks); - tabArea.performLayout(getLocalBounds()); + tabArea.performLayout(contentArea); positionErrorPopup(); } @@ -241,7 +198,9 @@ class ModelTab : public Component, private ChangeListener, public ChangeBroadcas int getMinimumRequiredHeightForWidth(int width) { - int height = 0; + width -= 2 * pagePadding; + + int height = 2 * pagePadding; height += modelInfoWidget.getPreferredHeightForWidth(width) + 2 * marginSize; @@ -270,28 +229,6 @@ class ModelTab : public Component, private ChangeListener, public ChangeBroadcas return height; } - void resetState() - { - model = std::make_shared(); - - // Publish the empty state so the status area does not keep - // showing the previous model's last status - model->setStatus(ModelStatus::EMPTY); - - modelSelectionWidget.resetState(); - modelInfoWidget.resetState(); - controlAreaWidget.resetState(); - inputTrackAreaWidget.resetState(); - outputTrackAreaWidget.resetState(); - - processCancelButton.setMode(processButtonInfo.displayLabel); - processCancelButton.setEnabled(false); - - currentProcessID = 0; - - resized(); - } - private: void initializeProcessCancelButton() { @@ -314,14 +251,6 @@ class ModelTab : public Component, private ChangeListener, public ChangeBroadcas addAndMakeVisible(processCancelButton); } - void changeListenerCallback(ChangeBroadcaster* source) - { - if (source == &modelSelectionWidget) - { - loadModelCallback(); - } - } - int getControlAreaRequiredHeightForTabWidth(int tabWidth) const { return jmax(minControlAreaHeight, controlAreaWidget.getRequiredHeightForWidth(tabWidth)); @@ -339,7 +268,7 @@ class ModelTab : public Component, private ChangeListener, public ChangeBroadcas } void addTrackSection(FlexBox& box, - Label& label, + Component& label, Component& trackArea, int numTracks, float totalTracks) const @@ -555,82 +484,53 @@ class ModelTab : public Component, private ChangeListener, public ChangeBroadcas URL(issueBaseUrl + query).launchInDefaultBrowser(); } - void loadModelCallback() + void finishLoading(const OpResult& result, const String& requestedPath) { - modelSelectionWidget.setDisabled(); - - // Disable processing until model is loaded - processCancelButton.setEnabled(false); - - // Obtain currently selected path - String selectedPath = modelSelectionWidget.getCurrentlySelectedPath(); - - DBG_AND_LOG("ModelTab::loadModelCallback: Attempting to load path \"" << selectedPath - << "\"."); - - loadingThreadPool.addJob( - [this, selectedPath] - { - OpResult result = model->load(selectedPath); - - // Perform updates on message (GUI) thread - MessageManager::callAsync( - [this, result] - { - if (abandoned) - { - // Tab was closed while this load was in flight - return; - } - - if (result.wasOk()) - { - modelSelectionWidget.setSuccessfulState(model->getLoadedPath()); - - modelInfoWidget.updateLabels(model->getMetadata()); - modelInfoWidget.addOpenablePath(model->getOpenablePath()); - - controlAreaWidget.updateControls(model->getControls()); + if (abandoned) + { + // Tab was closed while this load was in flight + return; + } - inputTrackAreaWidget.updateTracks(model->getInputTracks()); - outputTrackAreaWidget.updateTracks(model->getOutputTracks()); + if (result.wasOk()) + { + loading = false; - resized(); + catalog->recordLoadSuccess(model->getLoadedPath(), model->getMetadata()); - sendSynchronousChangeMessage(); + modelInfoWidget.updateLabels(model->getMetadata()); + modelInfoWidget.addOpenablePath(model->getOpenablePath()); - // Re-enable processing immediately - processCancelButton.setEnabled(true); + // Once loaded, only how the model is deployed is worth noting + if (const auto* entry = catalog->findEntry(model->getLoadedPath())) + modelInfoWidget.setBadges(ModelStyle::getBadges(*entry, false)); - notifyInitialLoadComplete(true); - } - else - { - const Error error = result.getError(); + controlAreaWidget.updateControls(model->getControls()); - std::function onExit = [this, error] - { - modelSelectionWidget.setUnsuccessfulState(error); + inputTrackAreaWidget.updateTracks(model->getInputTracks()); + outputTrackAreaWidget.updateTracks(model->getOutputTracks()); - // Re-enable processing after closing error window - processCancelButton.setEnabled(true); + resized(); - notifyInitialLoadComplete(false); - }; + // Enable processing now that a model is loaded + processCancelButton.setEnabled(true); - openErrorPopup(error, onExit); - } - }); - }); - } + sendSynchronousChangeMessage(); + } + else + { + const Error error = result.getError(); - void notifyInitialLoadComplete(bool wasSuccessful) - { - auto callback = std::move(initialLoadCallback); - initialLoadCallback = nullptr; + catalog->recordLoadFailure(requestedPath, error); - if (callback) - callback(this, wasSuccessful); + // The outcome is reported once the error has been seen + openErrorPopup(error, + [this] + { + loading = false; + sendChangeMessage(); + }); + } } void processCallback() @@ -679,7 +579,6 @@ class ModelTab : public Component, private ChangeListener, public ChangeBroadcas } } - modelSelectionWidget.setDisabled(); processCancelButton.setMode(cancelButtonInfo.displayLabel); // Switch choose-file button to inactive mode on all tracks during processing @@ -687,8 +586,10 @@ class ModelTab : public Component, private ChangeListener, public ChangeBroadcas uint64_t processID = currentProcessID; + SafePointer safeThis(this); + processingThreadPool.addJob( - [this, loadedInputFiles, processID] + [this, safeThis, loadedInputFiles, processID] { std::vector outputFiles; LabelList labels; @@ -711,48 +612,51 @@ class ModelTab : public Component, private ChangeListener, public ChangeBroadcas // Perform updates on message (GUI) thread MessageManager::callAsync( - [this, result, outputFilesPtr, labelsPtr] + [safeThis, result, outputFilesPtr, labelsPtr] { - if (abandoned) - { - // Tab was closed while this process was in flight - return; - } - - std::function onExit = [this] - { - // Re-enable processing immediately - modelSelectionWidget - .setFinishedState(); // TODO - should this be last selected? - processCancelButton.setMode(processButtonInfo.displayLabel); - - // Switch choose-file button back to active on all tracks - inputTrackAreaWidget.setLoadTrackEnabled(true); - }; - - if (result.wasOk()) - { - auto& outputMediaDisplays = outputTrackAreaWidget.getMediaDisplays(); - - for (size_t i = 0; - i < outputMediaDisplays.size() && i < outputFilesPtr->size(); - ++i) - { - outputMediaDisplays[i]->initializeDisplay( - URL((*outputFilesPtr)[i])); - outputMediaDisplays[i]->addLabels(*labelsPtr); - } - - onExit(); - } - else - { - openErrorPopup(result.getError(), onExit); - } + if (safeThis != nullptr) + safeThis->finishProcessing(result, *outputFilesPtr, *labelsPtr); }); }); } + void finishProcessing(const OpResult& result, + const std::vector& outputFiles, + const LabelList& labels) + { + if (abandoned) + { + // Tab was closed while this process was in flight + return; + } + + std::function onExit = [this] + { + // Re-enable processing immediately + processCancelButton.setMode(processButtonInfo.displayLabel); + + // Switch choose-file button back to active on all tracks + inputTrackAreaWidget.setLoadTrackEnabled(true); + }; + + if (result.wasOk()) + { + auto& outputMediaDisplays = outputTrackAreaWidget.getMediaDisplays(); + + for (size_t i = 0; i < outputMediaDisplays.size() && i < outputFiles.size(); ++i) + { + outputMediaDisplays[i]->initializeDisplay(URL(outputFiles[i])); + outputMediaDisplays[i]->addLabels(labels); + } + + onExit(); + } + else + { + openErrorPopup(result.getError(), onExit); + } + } + void cancelCallback() { processCancelButton.setEnabled(false); @@ -771,8 +675,6 @@ class ModelTab : public Component, private ChangeListener, public ChangeBroadcas } // Re-enable processing immediately - modelSelectionWidget.setFinishedState(); // TODO - should this be last selected? - processCancelButton.setMode(processButtonInfo.displayLabel); processCancelButton.setEnabled(true); @@ -927,27 +829,27 @@ class ModelTab : public Component, private ChangeListener, public ChangeBroadcas }; static constexpr float marginSize = 2; + // Space around the whole tab, matching the Home tab's + static constexpr int pagePadding = 8; - static constexpr int modelSelectionRowHeight = 30; static constexpr int minControlAreaHeight = 96; static constexpr int processButtonWidth = 150; static constexpr int processButtonRowHeight = 30; - static constexpr int trackSectionLabelHeight = 20; + static constexpr int trackSectionLabelHeight = ModelStyle::sectionHeaderHeight - 2; std::shared_ptr model { new Model() }; - ModelSelectionWidget modelSelectionWidget; ModelInfoWidget modelInfoWidget; ControlAreaWidget controlAreaWidget; - Label inputTracksLabel { "Input Tracks", "Input Tracks" }; + ModelStyle::SectionHeader inputTracksLabel { "Input Tracks" }; TrackAreaWidget inputTrackAreaWidget { DisplayMode::Input }; MultiButton processCancelButton; MultiButton::Mode processButtonInfo; MultiButton::Mode cancelButtonInfo; - Label outputTracksLabel { "Output Tracks", "Output Tracks" }; + ModelStyle::SectionHeader outputTracksLabel { "Output Tracks" }; TrackAreaWidget outputTrackAreaWidget { DisplayMode::Output }; ThreadPool loadingThreadPool { 1 }; @@ -955,7 +857,9 @@ class ModelTab : public Component, private ChangeListener, public ChangeBroadcas std::atomic currentProcessID { 0 }; std::atomic abandoned { false }; - std::function initialLoadCallback; + bool loading = false; + + SharedResourcePointer catalog; CentredAlertLookAndFeel centredAlertLF; std::unique_ptr errorPopupWindow; diff --git a/src/ModelTabContainer.h b/src/ModelTabContainer.h index 2d12604d..a701c757 100644 --- a/src/ModelTabContainer.h +++ b/src/ModelTabContainer.h @@ -1,27 +1,24 @@ /** - * @brief Adds tab container to HARP for MultiTabs + * @file ModelTabContainer.h + * @brief Tab bar holding the Home tab and one tab per opened model. * @author JEYuhas */ + #pragma once +#include + #include #include "HomeTab.h" -#include "Model.h" #include "ModelTab.h" -#include "widgets/ControlAreaWidget.h" -#include "widgets/ModelInfoWidget.h" -#include "widgets/ModelSelectionWidget.h" -#include "widgets/TrackAreaWidget.h" +#include "media/MediaDisplayComponent.h" -#include "utils/Errors.h" #include "utils/Interface.h" -#include "utils/Logging.h" -#include "utils/ModelRegistry.h" -#include "utils/Tutorial.h" +#include "utils/ModelCatalog.h" -#include "media/MediaDisplayComponent.h" +#include "windows/tutorial/TutorialTargets.h" using namespace juce; @@ -49,9 +46,9 @@ class ModelTabsLookAndFeel : public LookAndFeel_V4 const auto isActive = button.isFrontTab(); auto area = button.getActiveArea(); - const auto fill = isActive - ? activeTabColour - : inactiveTabColour.brighter(isMouseOver || isMouseDown ? 0.08f : 0.0f); + const auto fill = + isActive ? activeTabColour + : inactiveTabColour.brighter(isMouseOver || isMouseDown ? 0.08f : 0.0f); g.setColour(fill); g.fillRect(area); @@ -64,20 +61,29 @@ class ModelTabsLookAndFeel : public LookAndFeel_V4 auto textArea = button.getTextArea().reduced(tabTextInset, 0); - g.setColour(isActive ? activeTextColour - : inactiveTextColour); + g.setColour(isActive ? activeTextColour : inactiveTextColour); - g.drawText(button.getButtonText(), - textArea, - Justification::centred, - true); + g.drawText(button.getButtonText(), textArea, Justification::centred, true); } - int getTabButtonBestWidth(TabBarButton& button, int tabDepth) override + int getTabButtonBestWidth(TabBarButton& button, int /*tabDepth*/) override { - return button.getButtonText() == "Home" - ? homeTabWidth - : fixedTabWidth; + // The Home tab is always first + return button.getIndex() == 0 ? homeTabWidth : fixedTabWidth; + } + + /* The default gives a tab's close button the full height of the tab, which stretches it, + so it is kept square, centred, and clear of the tab's edge instead */ + Rectangle getTabButtonExtraComponentBounds(const TabBarButton& button, + Rectangle& textArea, + Component& extraComponent) override + { + textArea.removeFromRight(extraComponentMargin); + + const int side = extraComponent.getWidth(); + + return LookAndFeel_V4::getTabButtonExtraComponentBounds(button, textArea, extraComponent) + .withSizeKeepingCentre(side, side); } void drawTabButtonText(TabBarButton&, @@ -97,6 +103,7 @@ class ModelTabsLookAndFeel : public LookAndFeel_V4 static constexpr int fixedTabWidth = 140; static constexpr int homeTabWidth = 64; static constexpr int tabTextInset = 10; + static constexpr int extraComponentMargin = 6; }; /** @@ -156,19 +163,6 @@ class ModelTabPage : public Viewport Viewport::mouseWheelMove(e, wheel); } - /* Scrolling moves the tab under the tutorial overlay, which draws its highlight in - window coordinates and would otherwise keep pointing at where a component used - to be. */ - void visibleAreaChanged(const Rectangle&) override - { - if (onScrolled != nullptr) - { - onScrolled(); - } - } - - std::function onScrolled; - private: int layOutTabForVisibleWidth() { @@ -204,23 +198,29 @@ class ModelTabPage : public Viewport ModelTab& modelTab; }; -class ModelTabContainer : public TabbedComponent, - private ChangeListener, - public ChangeBroadcaster +class ModelTabContainer : public TabbedComponent, private ChangeListener, public ChangeBroadcaster { public: - ModelTabContainer() - : TabbedComponent(TabbedButtonBar::TabsAtTop) + ModelTabContainer() : TabbedComponent(TabbedButtonBar::TabsAtTop) { getTabbedButtonBar().setLookAndFeel(&tabsLookAndFeel); setColour(TabbedComponent::backgroundColourId, tabBackgroundColour); getTabbedButtonBar().setColour(TabbedButtonBar::tabTextColourId, Colours::white); getTabbedButtonBar().setColour(TabbedButtonBar::frontTextColourId, Colours::white); - getTabbedButtonBar().setColour(TabbedButtonBar::tabOutlineColourId, tabBackgroundColour.darker(0.35f)); - getTabbedButtonBar().setColour(TabbedButtonBar::frontOutlineColourId, tabBackgroundColour.darker(0.35f)); + getTabbedButtonBar().setColour(TabbedButtonBar::tabOutlineColourId, + tabBackgroundColour.darker(0.35f)); + getTabbedButtonBar().setColour(TabbedButtonBar::frontOutlineColourId, + tabBackgroundColour.darker(0.35f)); - createHomeTab(); + homeTab.onModelOpenRequested = [this](const String& modelPath, const String& modelName) + { openModelTab(modelPath, modelName); }; + + addTab("Home", tabBackgroundColour, &homeTab, false); + setCurrentTabIndex(0); + + setComponentID(TutorialTargets::modelTabs); + getTabbedButtonBar().setComponentID(TutorialTargets::tabBar); } ~ModelTabContainer() override @@ -248,32 +248,88 @@ class ModelTabContainer : public TabbedComponent, } } - ModelTab* createNewTab(const String& modelPath = {}, const String& modelName = {}) + /** + * Opens a model in a new tab and shows it. The tab fills in once the model has loaded, + * and closes by itself if loading fails, once the error has been dismissed. + */ + ModelTab* openModelTab(const String& modelPath, String tabName = {}) { - int index = getNumTabs(); + if (tabName.isEmpty()) + { + const auto* entry = catalog->findEntry(modelPath); + tabName = + entry != nullptr ? entry->name : modelPath.fromLastOccurrenceOf("/", false, false); + } auto* tab = new ModelTab(); // Owned from creation, so that a tab is never left unowned while a - // request is in flight (see closeModelTab and the destructor) + // request is in flight (see closeTab and the destructor) modelTabs.add(tab); + tab->addChangeListener(this); - auto tabName = modelName; - - if (tabName.isEmpty() && modelPath.isNotEmpty()) - tabName = ModelRegistry::getEntryForPath(modelPath).displayName; + auto* page = pages.add(new ModelTabPage(*tab)); - if (tabName.isEmpty()) - tabName = "Model " + String(index); + // The page is owned by pages and the tab by modelTabs, so it is added with + // deleteComponentWhenNotNeeded = false: closing it must not force JUCE + // to destroy the tab synchronously in removeTab(); see closeTab() for + // why destruction may need to be deferred. + addTab(tabName, tabBackgroundColour, page, false); + addCloseButton(*tab, getNumTabs() - 1); - addLoadedModelTab(tab, tabName); + setCurrentTabIndex(getNumTabs() - 1); - if (modelPath.isNotEmpty()) - tab->loadModelPath(modelPath); + tab->loadModel(modelPath); return tab; } + // Closes a model tab as if its close button had been clicked. Does nothing + // if the tab is no longer in the tab bar. + void closeTab(ModelTab* tabToClose) + { + const int index = findTabIndex(tabToClose); + + if (index < 0) + return; + + const auto currentIndex = getCurrentTabIndex(); + const auto targetIndex = currentIndex == index + ? jmax(0, index - 1) + : (currentIndex > index ? currentIndex - 1 : currentIndex); + + // Remove the tab from the UI immediately so it looks closed to the user. + // Because the page was added with deleteComponentWhenNotNeeded = false, + // removeTab() does not destroy it; it is owned via pages. + removeTab(index); + + // Deleting the page only detaches the tab from it + pages.removeObject(findPage(tabToClose)); + + setCurrentTabIndex(jlimit(0, getNumTabs() - 1, targetIndex)); + + tabToClose->removeChangeListener(this); + + if (tabToClose->hasPendingRequests()) + { + // A network request is still in flight. Destroying the tab + // (and its ThreadPool) now would force-kill a worker thread + // blocked in a network syscall. Abandoning it aborts the + // connection so the worker returns within moments; hand + // ownership to the reaper, which deletes it once it does. + tabToClose->abandon(); + + modelTabs.removeObject(tabToClose, false); + tabReaper.add(tabToClose); + } + else + { + modelTabs.removeObject(tabToClose, true); + } + + sendChangeMessage(); + } + ModelTabPage* getCurrentModelTabPage() const { return dynamic_cast(getCurrentContentComponent()); @@ -285,11 +341,6 @@ class ModelTabContainer : public TabbedComponent, return page != nullptr ? &page->getModelTab() : nullptr; } - HomeTab* getHomeTabIfShowing() const - { - return dynamic_cast(getCurrentContentComponent()); - } - void layOutCurrentPage() { if (auto* page = getCurrentModelTabPage()) @@ -304,158 +355,114 @@ class ModelTabContainer : public TabbedComponent, sendChangeMessage(); } - // Called whenever the page showing a model tab scrolls - std::function onPageScrolled; - - // Closes a model tab as if its close button had been clicked. Does nothing - // if the tab is no longer in the tab bar. - void closeTab(ModelTab* tab) { closeModelTab(tab); } - private: - void addLoadedModelTab(ModelTab* tab, const String& tabName) + // Close button shown on each model tab, which explains itself in the instructions box + struct CloseTabButton : public Button { - tab->addChangeListener(this); + CloseTabButton() : Button("Close") {} - auto* page = pages.add(new ModelTabPage(*tab)); - page->onScrolled = [this] + void paintButton(Graphics& g, bool isHighlighted, bool isDown) override { - if (onPageScrolled != nullptr) - onPageScrolled(); - }; + // Square whatever this is given, so that the circle and cross are never stretched + const float side = (float) jmin(getWidth(), getHeight()); + const auto bounds = + getLocalBounds().toFloat().withSizeKeepingCentre(side, side).reduced(1.0f); - // The page is owned by pages and the tab by modelTabs, so it is added with - // deleteComponentWhenNotNeeded = false: closing it must not force JUCE - // to destroy the tab synchronously in removeTab(); see closeModelTab() for - // why destruction may need to be deferred. - addTab(tabName, - tabBackgroundColour, - page, - false); + if (isHighlighted || isDown) + { + g.setColour(Colours::white.withAlpha(isDown ? 0.2f : 0.12f)); + g.fillEllipse(bounds); + } - addCloseButtonToModelTab(tab); + // A cross drawn from two strokes, which sits centred where a glyph might not + const auto cross = bounds.reduced(bounds.getWidth() * 0.32f); - setCurrentTabIndex(getNumTabs() - 1); - } + g.setColour(isHighlighted ? Colours::white : Colours::lightgrey); + g.drawLine({ cross.getTopLeft(), cross.getBottomRight() }, 1.5f); + g.drawLine({ cross.getBottomLeft(), cross.getTopRight() }, 1.5f); + } - void addCloseButtonToModelTab(ModelTab* tab) + void mouseEnter(const MouseEvent& e) override + { + Button::mouseEnter(e); + instructionsMessage->setMessage("Click to close this model tab."); + } + + void mouseExit(const MouseEvent& e) override + { + Button::mouseExit(e); + instructionsMessage->clearMessage(); + } + + SharedResourcePointer instructionsMessage; + }; + + void addCloseButton(ModelTab& tab, int tabIndex) { - auto* closeButton = new TextButton("x"); - closeButton->setTooltip("Close model tab"); + auto* closeButton = new CloseTabButton(); closeButton->setSize(18, 18); - closeButton->setColour(TextButton::buttonColourId, Colours::transparentBlack); - closeButton->setColour(TextButton::buttonOnColourId, Colours::transparentBlack); - closeButton->setColour(TextButton::textColourOffId, Colours::white); - closeButton->setColour(TextButton::textColourOnId, Colours::white); - closeButton->onClick = [this, tab] { closeModelTab(tab); }; + closeButton->onClick = [this, safeTab = SafePointer(&tab)] + { + // Deferred, since closing deletes the tab button that owns this one + MessageManager::callAsync( + [safeThis = SafePointer(this), safeTab] + { + if (safeThis != nullptr && safeTab != nullptr) + safeThis->closeTab(safeTab.getComponent()); + }); + }; - if (auto* tabButton = getTabbedButtonBar().getTabButton(getNumTabs() - 1)) + // The tab button takes ownership of the close button + if (auto* tabButton = getTabbedButtonBar().getTabButton(tabIndex)) tabButton->setExtraComponent(closeButton, TabBarButton::afterText); + else + delete closeButton; } - void closeModelTab(ModelTab* tabToClose) + ModelTabPage* findPage(const ModelTab* tab) const { - for (int i = 1; i < getNumTabs(); ++i) + for (auto* page : pages) { - auto* page = dynamic_cast(getTabContentComponent(i)); - - if (page != nullptr && &page->getModelTab() == tabToClose) - { - const auto currentIndex = getCurrentTabIndex(); - const auto targetIndex = currentIndex == i ? jmax(0, i - 1) - : (currentIndex > i ? currentIndex - 1 - : currentIndex); - - // Remove the tab from the UI immediately so it looks closed to - // the user. Because the tab was added with - // deleteComponentWhenNotNeeded = false, removeTab() does not - // destroy it; we own it via modelTabs. - removeTab(i); - - // Deleting the page only detaches the tab from it - pages.removeObject(page); - - if (getNumTabs() > 0) - setCurrentTabIndex(jlimit(0, getNumTabs() - 1, targetIndex)); - - tabToClose->removeChangeListener(this); - - if (tabToClose->hasPendingRequests()) - { - // A network request is still in flight. Destroying the tab - // (and its ThreadPool) now would force-kill a worker thread - // blocked in a network syscall. Abandoning it aborts the - // connection so the worker returns within moments; hand - // ownership to the reaper, which deletes it once it does. - tabToClose->abandon(); - - modelTabs.removeObject(tabToClose, false); - tabReaper.add(tabToClose); - } - else - { - modelTabs.removeObject(tabToClose, true); - } - - sendChangeMessage(); - return; - } + if (&page->getModelTab() == tab) + return page; } + + return nullptr; } - void createHomeTab() + int findTabIndex(const ModelTab* tab) const { - auto* homeTab = new HomeTab(); - homeTab->onModelLoadRequested = [this, homeTab](String modelPath, String modelName) + for (int i = 1; i < getNumTabs(); ++i) { - // Owned from creation as well: this tab is not in the UI yet, but it - // has a load request in flight, so quitting the app now must find it - // in modelTabs and abandon it rather than leave it running. - auto* pendingTab = new ModelTab(); - modelTabs.add(pendingTab); - - pendingTab->onNextModelLoadComplete( - [this, homeTab, modelName](ModelTab* tab, bool wasSuccessful) - { - if (wasSuccessful) - { - addLoadedModelTab(tab, modelName); - sendChangeMessage(); - } - else - { - // Deleted asynchronously because this runs from inside - // the tab's own load completion handler - modelTabs.removeObject(tab, false); - MessageManager::callAsync([tab] { delete tab; }); - } - - homeTab->resetSelection(); - }); - - pendingTab->loadModelPath(modelPath); - }; + auto* page = dynamic_cast(getTabContentComponent(i)); - addTab("Home", - tabBackgroundColour, - homeTab, - false); + if (page != nullptr && &page->getModelTab() == tab) + return i; + } - setCurrentTabIndex(0); + return -1; } void changeListenerCallback(ChangeBroadcaster* source) override { - if (auto* tab = dynamic_cast(source)) - { - // What the tab has to show changed, and so how far its page has to scroll - for (auto* page : pages) - { - if (&page->getModelTab() == tab) - page->layOutTab(); - } + auto* tab = dynamic_cast(source); - sendChangeMessage(); // bubble up to MainComponent + if (tab == nullptr) + return; + + // A tab left without a model once its load finished failed to load one, and + // has nothing to show now that its error has been dismissed + if (! tab->isModelLoaded() && ! tab->isLoading()) + { + closeTab(tab); + return; } + + // What the tab has to show changed, and so how far its page has to scroll + if (auto* page = findPage(tab)) + page->layOutTab(); + + sendChangeMessage(); // bubble up to MainComponent } // Owns model tabs whose UI has been closed while a request was still in @@ -512,14 +519,17 @@ class ModelTabContainer : public TabbedComponent, OwnedArray pendingTabs; }; - const Colour tabBackgroundColour { - getUIColourIfAvailable(LookAndFeel_V4::ColourScheme::UIColour::windowBackground) - }; + const Colour tabBackgroundColour { getUIColourIfAvailable( + LookAndFeel_V4::ColourScheme::UIColour::windowBackground) }; ModelTabsLookAndFeel tabsLookAndFeel; + HomeTab homeTab; + OwnedArray modelTabs; // Declared after modelTabs so that each page is destroyed before the tab it shows OwnedArray pages; DeferredTabReaper tabReaper; + + SharedResourcePointer catalog; }; diff --git a/src/clients/Client.h b/src/clients/Client.h index 7ecf154a..49dc5f5e 100644 --- a/src/clients/Client.h +++ b/src/clients/Client.h @@ -255,6 +255,80 @@ class RegisteredWebInputStream : public WebInputStream RequestRegistry& registry; }; +/* + Equivalent to URL::createInputStream(), except that the stream is + registered with requestRegistry for its whole lifetime so that it can be + aborted locally (see RequestRegistry). Registering before connecting is + what makes the connect phase abortable too - that phase blocks for up to + the connection timeout, which is two minutes for process requests. + + Returns nullptr if the connection could not be established, including + when the request was aborted. Without a registry (e.g. for the short-lived + client used for token validation) this behaves exactly as before. +*/ +inline std::unique_ptr createRegisteredStream(RequestRegistry* requestRegistry, + const URL& endpoint, + const URL::InputStreamOptions& options) +{ + if (requestRegistry == nullptr || endpoint.isLocalFile()) + { + return endpoint.createInputStream(options); + } + + if (requestRegistry->hasBeenAborted()) + { + // Requests were aborted, so do not open another connection + return nullptr; + } + + auto stream = std::make_unique( + *requestRegistry, + endpoint, + options.getParameterHandling() == URL::ParameterHandling::inPostData); + + const String extraHeaders = options.getExtraHeaders(); + + if (extraHeaders.isNotEmpty()) + { + stream->withExtraHeaders(extraHeaders); + } + + const int connectionTimeoutMs = options.getConnectionTimeoutMs(); + + if (connectionTimeoutMs != 0) + { + stream->withConnectionTimeout(connectionTimeoutMs); + } + + const String requestCmd = options.getHttpRequestCmd(); + + if (requestCmd.isNotEmpty()) + { + stream->withCustomRequestCommand(requestCmd); + } + + stream->withNumRedirectsToFollow(options.getNumRedirectsToFollow()); + + const bool connected = stream->connect(nullptr); + + if (int* statusCode = options.getStatusCode()) + { + *statusCode = stream->getStatusCode(); + } + + if (StringPairArray* responseHeaders = options.getResponseHeaders()) + { + *responseHeaders = stream->getResponseHeaders(); + } + + if (! connected || stream->isError()) + { + return nullptr; + } + + return stream; +} + class Client { public: @@ -292,77 +366,11 @@ class Client virtual String inferEndpointPath(String modelPath) = 0; virtual String inferDocumentationPath(String modelPath) = 0; - /* - Equivalent to URL::createInputStream(), except that the stream is - registered with requestRegistry for its whole lifetime so that it can be - aborted locally (see RequestRegistry). Registering before connecting is - what makes the connect phase abortable too - that phase blocks for up to - the connection timeout, which is two minutes for process requests. - - Returns nullptr if the connection could not be established, including - when the request was aborted. Clients without a registry (e.g. the - short-lived one used for token validation) behave exactly as before. - */ + // See createRegisteredStream std::unique_ptr createRequestStream(const URL& endpoint, const URL::InputStreamOptions& options) const { - if (requestRegistry == nullptr || endpoint.isLocalFile()) - { - return endpoint.createInputStream(options); - } - - if (requestRegistry->hasBeenAborted()) - { - // Requests were aborted, so do not open another connection - return nullptr; - } - - auto stream = std::make_unique( - *requestRegistry, - endpoint, - options.getParameterHandling() == URL::ParameterHandling::inPostData); - - const String extraHeaders = options.getExtraHeaders(); - - if (extraHeaders.isNotEmpty()) - { - stream->withExtraHeaders(extraHeaders); - } - - const int connectionTimeoutMs = options.getConnectionTimeoutMs(); - - if (connectionTimeoutMs != 0) - { - stream->withConnectionTimeout(connectionTimeoutMs); - } - - const String requestCmd = options.getHttpRequestCmd(); - - if (requestCmd.isNotEmpty()) - { - stream->withCustomRequestCommand(requestCmd); - } - - stream->withNumRedirectsToFollow(options.getNumRedirectsToFollow()); - - const bool connected = stream->connect(nullptr); - - if (int* statusCode = options.getStatusCode()) - { - *statusCode = stream->getStatusCode(); - } - - if (StringPairArray* responseHeaders = options.getResponseHeaders()) - { - *responseHeaders = stream->getResponseHeaders(); - } - - if (! connected || stream->isError()) - { - return nullptr; - } - - return stream; + return createRegisteredStream(requestRegistry, endpoint, options); } OpResult queryToken(const String& tokenToQuery, String& response, const int timeoutMs = 10000) diff --git a/src/clients/providers/stability/StabilityClient.h b/src/clients/providers/stability/StabilityClient.h index 9652bba6..e7fcb1ad 100644 --- a/src/clients/providers/stability/StabilityClient.h +++ b/src/clients/providers/stability/StabilityClient.h @@ -78,6 +78,12 @@ class StabilityClient : public Client return isValidTextToAudioPath(modelPath) || isValidAudioToAudioPath(modelPath); } + // Every model this provider offers, each with built-in controls + static StringArray getModelPaths() + { + return { "stability/text-to-audio", "stability/audio-to-audio" }; + } + String inferHostSlashModel(String modelPath) override { String hostSlashModel; diff --git a/src/clients/providers/stability/models/audio-to-audio.json b/src/clients/providers/stability/models/audio-to-audio.json index d46a92d0..b78d691e 100644 --- a/src/clients/providers/stability/models/audio-to-audio.json +++ b/src/clients/providers/stability/models/audio-to-audio.json @@ -4,7 +4,14 @@ "description": "Enter text prompt to modify an audio file.", "author": "Stability AI", "tags": [ - "audio-to-audio", + "category:generation", + "subcategory:editing", + "input:audio", + "input:text", + "output:audio/wav|mp3", + "sample-rate:44100", + "channels:stereo", + "stable-audio", "stability" ] }, @@ -16,6 +23,7 @@ }, { "label": "Duration (s)", + "info": "Length of the generated audio, in seconds.", "minimum": 1, "maximum": 190, "step": 1, @@ -24,6 +32,7 @@ }, { "label": "steps", + "info": "Number of diffusion steps. Higher values improve quality but take longer to process.", "minimum": 30, "maximum": 100, "step": 1, @@ -32,6 +41,7 @@ }, { "label": "cfg", + "info": "Classifier-free guidance. Higher values make the output follow the prompt more strictly.", "minimum": 1, "maximum": 25, "step": 1, @@ -40,6 +50,7 @@ }, { "label": "Output Format", + "info": "File format of the generated audio.", "choices": [ [ "wav", @@ -55,6 +66,7 @@ }, { "label": "Text Prompt", + "info": "Description of how to transform the input audio.", "value": "A song in the 3/4 time signature that features cheerful acoustic guitar, live recorded drums, and rhythmic claps, The mood is happy and up-lifting.", "type": "text_box" } diff --git a/src/clients/providers/stability/models/text-to-audio.json b/src/clients/providers/stability/models/text-to-audio.json index 49f44d01..2a207ff0 100644 --- a/src/clients/providers/stability/models/text-to-audio.json +++ b/src/clients/providers/stability/models/text-to-audio.json @@ -4,13 +4,20 @@ "description": "Enter text prompts to generate 44.1khz audio.", "author": "Stability AI", "tags": [ - "text-to-audio", + "category:generation", + "subcategory:holistic", + "input:text", + "output:audio/wav|mp3", + "sample-rate:44100", + "channels:stereo", + "stable-audio", "stability" ] }, "inputs": [ { "label": "Duration (s)", + "info": "Length of the generated audio, in seconds.", "minimum": 1, "maximum": 190, "step": 1, @@ -19,6 +26,7 @@ }, { "label": "steps", + "info": "Number of diffusion steps. Higher values improve quality but take longer to process.", "minimum": 30, "maximum": 100, "step": 1, @@ -27,6 +35,7 @@ }, { "label": "cfg", + "info": "Classifier-free guidance. Higher values make the output follow the prompt more strictly.", "minimum": 1, "maximum": 25, "step": 1, @@ -35,6 +44,7 @@ }, { "label": "Output Format", + "info": "File format of the generated audio.", "choices": [ [ "wav", @@ -50,6 +60,7 @@ }, { "label": "Text Prompt", + "info": "Description of the audio to generate.", "value": "A song in the 3/4 time signature that features cheerful acoustic guitar, live recorded drums, and rhythmic claps, The mood is happy and up-lifting.", "type": "text_box" } diff --git a/src/gui/HoverableLabel.h b/src/gui/HoverableLabel.h index 8c722af2..a467339f 100644 --- a/src/gui/HoverableLabel.h +++ b/src/gui/HoverableLabel.h @@ -99,20 +99,16 @@ class HoverableLabel : public Label std::function onClick; private: + /* Where the text is drawn, in this label's own coordinates (which are what hitTest is + given), allowing for its border and justification */ Rectangle getTextBounds() const { - Font f = getFont(); + const Font f = getFont(); + const Rectangle textSize(GlyphArrangement::getStringWidthInt(f, getText()), + roundToInt(f.getHeight())); - int textWidth = f.getStringWidth(getText()); - int textHeight = f.getHeight(); - - float x_offset = (getBounds().getWidth() - textWidth) / 2; - float y_offset = (getBounds().getHeight() - textHeight) / 2; - - return Rectangle(getX() + static_cast(x_offset), - getY() + static_cast(y_offset), - static_cast(textWidth), - static_cast(textHeight)); + return getJustificationType().appliedToRectangle( + textSize, getBorderSize().subtractedFrom(getLocalBounds())); } bool hoverable; diff --git a/src/media/MediaDisplayComponent.cpp b/src/media/MediaDisplayComponent.cpp index 67b952f3..c0c892e0 100644 --- a/src/media/MediaDisplayComponent.cpp +++ b/src/media/MediaDisplayComponent.cpp @@ -4,6 +4,8 @@ #include "../utils/Interface.h" +#include "../windows/tutorial/TutorialTargets.h" + #include void OptionalBannerComponent::paint(Graphics& g) @@ -206,6 +208,9 @@ MediaDisplayComponent::MediaDisplayComponent(String name, bool req, bool fromDAW void MediaDisplayComponent::initializeButtons() { + playStopButton.setComponentID(TutorialTargets::trackPlayButton); + chooseFileButton.setComponentID(TutorialTargets::trackFolderButton); + // Mode when a playable file is loaded playButtonActiveInfo = MultiButton::Mode { "Play-Active", "Click to start playback.", @@ -1298,26 +1303,6 @@ void MediaDisplayComponent::updateCursorPosition() Rectangle(cursorPositionX, cursorPositionY, cursorWidth, mediaBounds.getHeight())); } -Rectangle MediaDisplayComponent::getChooseFileButtonBounds() -{ - if (auto* p = chooseFileButton.getParentComponent()) - { - return getLocalArea(p, chooseFileButton.getBounds()); - } - - return chooseFileButton.getBounds(); -} - -Rectangle MediaDisplayComponent::getPlayButtonBounds() -{ - if (auto* p = playStopButton.getParentComponent()) - { - return getLocalArea(p, playStopButton.getBounds()); - } - - return playStopButton.getBounds(); -} - void MediaDisplayComponent::mouseEnter(const MouseEvent& e) { if (! isThumbnailTrack() && e.eventComponent == getMediaComponent() diff --git a/src/media/MediaDisplayComponent.h b/src/media/MediaDisplayComponent.h index 7e6dc279..c98f87a3 100644 --- a/src/media/MediaDisplayComponent.h +++ b/src/media/MediaDisplayComponent.h @@ -155,8 +155,6 @@ class MediaDisplayComponent : public Component, virtual bool isPlaying() { return transportSource.isPlaying(); } - Rectangle getChooseFileButtonBounds(); - Rectangle getPlayButtonBounds(); int getNumOverheadLabels(); diff --git a/src/utils/Clients.h b/src/utils/Clients.h index 7b96759d..605ca1c8 100644 --- a/src/utils/Clients.h +++ b/src/utils/Clients.h @@ -83,3 +83,38 @@ inline OpResult multiplexClients(String modelPath, std::unique_ptr& clie return OpResult::ok(); } + +/** + * Whether any provider recognizes a model path, so that one no client could load can be + * rejected before anything is attempted with it. + */ +inline bool isSupportedModelPath(const String& modelPath) +{ + return StabilityClient::matchesPathSpec(modelPath) || GradioClient::matchesPathSpec(modelPath); +} + +/** + * The model cards of models whose controls ship with HARP, keyed by model path, so that + * they can be listed without querying anything. + */ +inline std::vector> getBuiltInModelCards() +{ + static const Identifier cardKey { "card" }; + + std::vector> cards; + + StabilityClient stabilityClient; + + for (const auto& modelPath : StabilityClient::getModelPaths()) + { + DynamicObject::Ptr controls; + + if (stabilityClient.queryControls(modelPath, controls).wasOk() && controls != nullptr) + { + if (auto* card = controls->getProperty(cardKey).getDynamicObject()) + cards.emplace_back(modelPath, card); + } + } + + return cards; +} diff --git a/src/utils/ModelCatalog.h b/src/utils/ModelCatalog.h new file mode 100644 index 00000000..f83a67c4 --- /dev/null +++ b/src/utils/ModelCatalog.h @@ -0,0 +1,593 @@ +/** + * @file ModelCatalog.h + * @brief The models HARP offers for browsing on the Home tab. + * + * The catalog combines the models whose controls ship with HARP (Stability AI), every + * Space of the HARP organization on Hugging Face, and any other paths the user has loaded + * successfully. The organization's Spaces are listed through the Hub API in a single + * request, which describes each Space (including its tags, see ModelTags.h) without waking + * any of them. The last listing is cached on disk, so the Home tab is populated at once on + * startup and still works offline. + */ + +#pragma once + +#include +#include +#include +#include + +#include + +#include "../Model.h" + +#include "Clients.h" +#include "Errors.h" +#include "Logging.h" +#include "ModelTags.h" +#include "Settings.h" + +using namespace juce; + +struct CatalogEntry +{ + // Outcome of the last attempt to load a model in this session + enum class LoadStatus + { + None, // Not attempted, or loaded successfully + Failed, // Failed for a reason specific to the model + Unavailable, // The Space is not running (e.g., it crashed or was paused) + TryAgain // Failed for a reason that is likely temporary + }; + + // A short note on whether the model can be expected to load + struct AvailabilityNote + { + String text; // Empty if there is nothing to note + bool isProblem = false; // As opposed to merely informative, e.g. "Sleeping" + }; + + String path; + String name; + String description; + String provider; + ModelTags tags; + + // Deployment details reported by the Hub, empty when unknown + String stage; // e.g. "RUNNING", "SLEEPING", "RUNTIME_ERROR" + String hardware; // e.g. "cpu-basic", "zero-a10g" + + bool isCustom = false; // Added by the user rather than found + LoadStatus loadStatus = LoadStatus::None; + + bool isZeroGPU() const { return hardware.startsWithIgnoreCase("zero-"); } + + /* The Hub's report on a Space, before anything is asked of it, is only an expectation. + Once a load has been attempted, its outcome is what counts. */ + AvailabilityNote getAvailabilityNote() const + { + switch (loadStatus) + { + case LoadStatus::Failed: // Only custom paths are still listed after this + return { "Failed to load", true }; + case LoadStatus::Unavailable: + return { "Unavailable", true }; + case LoadStatus::TryAgain: + return { "Try again", true }; + case LoadStatus::None: + break; + } + + if (stage == "SLEEPING") + return { "Sleeping" }; // Loads, after a wait while the Space starts + if (stage.contains("BUILDING") || stage.contains("STARTING")) + return { "Starting" }; + + return {}; + } + + // Whether the Hub reports the Space as unable to serve at all, e.g. crashed or paused + bool isReportedBroken() const + { + return stage.contains("ERROR") || stage == "PAUSED" || stage == "STOPPED" + || stage == "NO_APP_FILE" || stage == "DELETING"; + } + + String getSearchableText() const + { + return (name + " " + description + " " + path + " " + provider + " " + + tags.getSearchableText()) + .toLowerCase(); + } +}; + +class ModelCatalog : public ChangeBroadcaster +{ +public: + // Hugging Face organization whose Spaces are listed + static constexpr const char* hubOrganization = "teamup-tech"; + + enum class FetchState + { + Fetching, + Succeeded, + Failed + }; + + ModelCatalog() + { + for (const auto& [path, card] : getBuiltInModelCards()) + builtInEntries.push_back(makeEntry(path, ModelMetadata(card.get()), "Stability AI")); + + customPaths = StringArray::fromLines(Settings::getString(customPathsKey)); + customPaths.removeEmptyStrings(); + + loadCachedListing(); + rebuildEntries(); + + refresh(); + } + + ~ModelCatalog() override + { + /* A listing in flight would leave its worker blocked in a network call, and + destroying the pool under it would force-kill the thread. Abort the request so + the worker returns at once, then wait for it. */ + requestRegistry.abortActiveRequests(); + fetchPool.removeAllJobs(true, 5000); + } + + // All entries, built-in models first, then the organization's, then custom paths + const std::vector& getEntries() const { return entries; } + + const CatalogEntry* findEntry(const String& path) const + { + for (const auto& entry : entries) + { + if (entry.path.equalsIgnoreCase(path)) + return &entry; + } + + return nullptr; + } + + // Lists the organization's Spaces again, in the background + void refresh() + { + if (fetchState == FetchState::Fetching && fetchPool.getNumJobs() > 0) + return; + + fetchState = FetchState::Fetching; + sendChangeMessage(); + + WeakReference weakThis(this); + + fetchPool.addJob( + [this, weakThis] + { + auto listing = std::make_shared(); + String error; + + const bool succeeded = fetchHubListing(*listing, error); + + MessageManager::callAsync( + [weakThis, succeeded, listing, error] + { + if (auto* catalog = weakThis.get()) + catalog->finishRefresh(succeeded, *listing, error); + }); + }); + } + + FetchState getFetchState() const { return fetchState; } + String getFetchError() const { return fetchError; } + // When the listing being shown was obtained, or a null time if there is none + Time getListingTime() const { return listingTime; } + int getNumHubEntries() const { return (int) (hubEntries.size() - hiddenEntries.size()); } + + // Spaces left out because they cannot currently be loaded, with the reason for each + struct HiddenEntry + { + String name; + String reason; + }; + + const std::vector& getHiddenEntries() const { return hiddenEntries; } + + static String getHubOrganizationURL() + { + return "https://huggingface.co/" + String(hubOrganization); + } + + /* A loaded model's own card is more authoritative than what its listing says about it, + so it is shown in the listing's place for the rest of the session. */ + void recordLoadSuccess(const String& path, const ModelMetadata& card) + { + loadStatuses.erase(path.toLowerCase()); + loadedCards[path.toLowerCase()] = card; + + if (findEntry(path) == nullptr) + { + customPaths.add(path); + saveCustomPaths(); + } + + rebuildEntries(); + sendChangeMessage(); + } + + void recordLoadFailure(const String& path, const Error& error) + { + // Paths that never loaded are not listed, so a typo does not become an entry + if (findEntry(path) == nullptr) + return; + + loadStatuses[path.toLowerCase()] = classifyLoadFailure(error); + + rebuildEntries(); + sendChangeMessage(); + } + + void removeCustomPath(const String& path) + { + customPaths.removeString(path, true); + saveCustomPaths(); + + rebuildEntries(); + sendChangeMessage(); + } + + /* What a failed load says about the model: + - TryAgain: nothing about the model itself, e.g. the connection dropped, the request + timed out or was rate limited, the GPU quota ran out, or the Space was still + waking up. The same load is likely to work later. + - Unavailable: the Space is not running, e.g. it crashed or was paused. + - Failed: the Space answered, but not as a HARP model, e.g. it does not exist, is + private, or its interface could not be read. Retrying will not help. */ + static CatalogEntry::LoadStatus classifyLoadFailure(const Error& error) + { + using LoadStatus = CatalogEntry::LoadStatus; + + if (const auto* httpError = std::get_if(&error)) + { + if (httpError->type == HttpError::Type::ConnectionFailed) + return LoadStatus::TryAgain; + + if (httpError->type == HttpError::Type::BadStatusCode && httpError->statusCode == 429) + return LoadStatus::TryAgain; + + // Only reaches here when the Hub could not say why (see GradioClient) + if (httpError->type == HttpError::Type::BadStatusCode && httpError->statusCode == 503) + return LoadStatus::Unavailable; + } + else if (const auto* gradioError = std::get_if(&error)) + { + switch (gradioError->type) + { + case GradioError::Type::SpaceStarting: + case GradioError::Type::IncompleteResponse: + case GradioError::Type::QuotaExceeded: + return LoadStatus::TryAgain; + + case GradioError::Type::SpaceUnavailable: + return LoadStatus::Unavailable; + + case GradioError::Type::RuntimeError: + case GradioError::Type::Indeterminate: + break; + } + } + + return LoadStatus::Failed; + } + +private: + static CatalogEntry + makeEntry(const String& path, const ModelMetadata& metadata, const String& provider) + { + CatalogEntry entry; + entry.path = path; + entry.name = metadata.name.empty() ? getNameFromPath(path) : String(metadata.name); + entry.description = metadata.description; + entry.provider = provider; + + StringArray tags; + + for (const auto& tag : metadata.tags) + tags.add(tag); + + entry.tags = ModelTags::parse(tags); + + return entry; + } + + static String getNameFromPath(const String& path) + { + return path.fromLastOccurrenceOf("/", false, false).replaceCharacters("-_", " ").trim(); + } + + /* The Hub pages its results, so the request is repeated for as long as the response + links to a next page. Runs on the fetch thread. */ + bool fetchHubListing(var& listing, String& error) + { + static constexpr int timeoutMs = 15000; + static constexpr int maxPages = 20; + + URL pageURL = URL("https://huggingface.co/api/spaces") + .withParameter("author", hubOrganization) + .withParameter("limit", "1000") + .withParameter("expand[]", "cardData") + .withParameter("expand[]", "runtime") + .withParameter("expand[]", "private") + .withParameter("expand[]", "disabled"); + + Array spaces; + + for (int page = 0; page < maxPages && pageURL.isWellFormed(); ++page) + { + DBG_AND_LOG("ModelCatalog::fetchHubListing: Requesting \"" << pageURL.toString(true) + << "\"."); + + int statusCode = 0; + StringPairArray responseHeaders; + + auto options = URL::InputStreamOptions(URL::ParameterHandling::inAddress) + .withConnectionTimeoutMs(timeoutMs) + .withStatusCode(&statusCode) + .withResponseHeaders(&responseHeaders) + .withNumRedirectsToFollow(5); + + auto stream = createRegisteredStream(&requestRegistry, pageURL, options); + + if (stream == nullptr) + { + error = requestRegistry.hasBeenAborted() ? "the request was aborted" + : "could not connect"; + return false; + } + + if (statusCode != 200) + { + error = "status code " + String(statusCode); + return false; + } + + const var response = JSON::parse(stream->readEntireStreamAsString()); + + if (! response.isArray()) + { + error = "unexpected response"; + return false; + } + + spaces.addArray(*response.getArray()); + + pageURL = URL(getNextPageLink(responseHeaders)); + } + + listing = spaces; + + DBG_AND_LOG("ModelCatalog::fetchHubListing: Found " << spaces.size() << " Spaces."); + + return true; + } + + // The target of a 'Link: <...>; rel="next"' header, or empty if there is none + static String getNextPageLink(const StringPairArray& headers) + { + for (const auto& link : StringArray::fromTokens(headers["link"], ",", "<>")) + { + if (link.contains("rel=\"next\"")) + return link.fromFirstOccurrenceOf("<", false, false) + .upToFirstOccurrenceOf(">", false, false) + .trim(); + } + + return {}; + } + + void finishRefresh(bool succeeded, const var& listing, const String& error) + { + if (succeeded) + { + hubEntries = parseHubListing(listing, true); + listingTime = Time::getCurrentTime(); + fetchState = FetchState::Succeeded; + fetchError.clear(); + + saveCachedListing(listing); + } + else + { + DBG_AND_LOG("ModelCatalog::finishRefresh: Could not list models (" << error << ")."); + + fetchState = FetchState::Failed; + fetchError = error; + } + + rebuildEntries(); + sendChangeMessage(); + } + + std::vector parseHubListing(const var& listing, bool isCurrent) const + { + std::vector parsed; + + if (auto* spaces = listing.getArray()) + { + for (const auto& space : *spaces) + { + const String path = space["id"].toString(); + + // Hidden from anyone without access, or switched off by the Hub + if (path.isEmpty() || (bool) space["private"] || (bool) space["disabled"]) + continue; + + const var& cardData = space["cardData"]; + + StringArray tags; + + if (auto* tagList = cardData["tags"].getArray()) + { + for (const auto& tag : *tagList) + tags.add(tag.toString()); + } + + CatalogEntry entry; + entry.path = path; + entry.name = cardData["title"].toString().trim(); + // A summary of the model card's description, which the Space's README is + // recommended to carry (see "Listing on the Home Tab" in the pyharp README). + // The card itself replaces it once the model is loaded. + entry.description = cardData["short_description"].toString().trim(); + entry.provider = "Hugging Face"; + entry.tags = ModelTags::parse(tags); + entry.hardware = space["runtime"]["hardware"]["current"].toString(); + + // A cached stage is stale, and would misreport whether the Space is awake + if (isCurrent) + entry.stage = space["runtime"]["stage"].toString().toUpperCase(); + + if (entry.name.isEmpty()) + entry.name = getNameFromPath(path); + + parsed.push_back(std::move(entry)); + } + } + + std::sort(parsed.begin(), + parsed.end(), + [](const CatalogEntry& a, const CatalogEntry& b) + { return a.name.compareNatural(b.name) < 0; }); + + return parsed; + } + + void rebuildEntries() + { + entries = builtInEntries; + hiddenEntries.clear(); + + /* A Space that cannot load is not offered, whether the Hub reports it as broken or + it answered a load this session in a way no retry will change (e.g., it is not a + HARP app, or needs controls this version does not support). Custom paths are the + user's own, so they stay listed with a note instead. */ + for (const auto& entry : hubEntries) + { + auto status = loadStatuses.find(entry.path.toLowerCase()); + + const bool failedToLoad = status != loadStatuses.end() + && status->second == CatalogEntry::LoadStatus::Failed; + + if (failedToLoad) + hiddenEntries.push_back({ entry.name, "failed to load" }); + else if (entry.isReportedBroken()) + hiddenEntries.push_back( + { entry.name, "the Space reports " + entry.stage.replace("_", " ").toLowerCase() }); + else + entries.push_back(entry); + } + + for (const auto& path : customPaths) + { + if (findEntry(path) != nullptr) + continue; + + CatalogEntry entry; + entry.path = path; + entry.name = getNameFromPath(path); + entry.provider = "Custom"; + entry.isCustom = true; + + entries.push_back(std::move(entry)); + } + + for (auto& entry : entries) + { + const String key = entry.path.toLowerCase(); + + auto status = loadStatuses.find(key); + + entry.loadStatus = + status != loadStatuses.end() ? status->second : CatalogEntry::LoadStatus::None; + + if (auto loaded = loadedCards.find(key); loaded != loadedCards.end()) + { + const ModelMetadata& card = loaded->second; + + if (! card.description.empty()) + entry.description = card.description; + + if (entry.isCustom && ! card.name.empty()) + entry.name = card.name; + + if (! card.tags.empty()) + { + StringArray tags; + + for (const auto& tag : card.tags) + tags.add(tag); + + entry.tags = ModelTags::parse(tags); + } + } + } + } + + static File getCacheFile() + { + if (auto* settings = Settings::getUserSettings()) + return settings->getFile().getSiblingFile("model_catalog.json"); + + return {}; + } + + void loadCachedListing() + { + const File cacheFile = getCacheFile(); + + if (! cacheFile.existsAsFile()) + return; + + hubEntries = parseHubListing(JSON::parse(cacheFile), false); + listingTime = cacheFile.getLastModificationTime(); + } + + static void saveCachedListing(const var& listing) + { + const File cacheFile = getCacheFile(); + + if (cacheFile != File() && ! cacheFile.replaceWithText(JSON::toString(listing, true))) + { + DBG_AND_LOG("ModelCatalog::saveCachedListing: Could not write \"" + << cacheFile.getFullPathName() << "\"."); + } + } + + void saveCustomPaths() + { + Settings::setValue(customPathsKey, customPaths.joinIntoString("\n"), true); + } + + static constexpr const char* customPathsKey = "models.customPaths"; + + std::vector builtInEntries; + std::vector hubEntries; + StringArray customPaths; + // Keyed by lowercase path, since the Hub does not distinguish case + std::map loadStatuses; + std::map loadedCards; + std::vector hiddenEntries; + + std::vector entries; + + FetchState fetchState = FetchState::Failed; + String fetchError; + Time listingTime; + + RequestRegistry requestRegistry; + // Declared last so that it is destroyed first, while what its job uses still exists + ThreadPool fetchPool { 1 }; + + JUCE_DECLARE_WEAK_REFERENCEABLE(ModelCatalog) +}; diff --git a/src/utils/ModelRegistry.h b/src/utils/ModelRegistry.h deleted file mode 100644 index 71e56503..00000000 --- a/src/utils/ModelRegistry.h +++ /dev/null @@ -1,156 +0,0 @@ -/** - * @file ModelRegistry.h - * @brief Temporary model registry accessors for model discovery. - */ - -#pragma once - -#include - -#include - -using namespace juce; - -namespace ModelRegistry -{ -struct Entry -{ - String path; - String displayName; - String summary; - String provider; - std::vector tags; - - Entry() = default; - Entry(String p, String dn, String s, String pr, std::vector t = {}) - : path(std::move(p)), displayName(std::move(dn)), summary(std::move(s)), provider(std::move(pr)), tags(std::move(t)) - {} -}; - -inline String getCleanModelPath(const String& modelPath) -{ - auto cleaned = modelPath.trim(); - - for (const auto& tag : { String(" [ERROR]"), String(" [DOWN]"), String(" [TRY AGAIN]"), String(" [SLEEPING]") }) - cleaned = cleaned.replace(tag, ""); - - return cleaned.trim(); -} - -inline String getFallbackModelDisplayName(const String& modelPath) -{ - auto cleaned = getCleanModelPath(modelPath).upToFirstOccurrenceOf(" [", false, false).trim(); - auto tokens = StringArray::fromTokens(cleaned, "/", ""); - - if (tokens.size() > 0) - return tokens[tokens.size() - 1].replaceCharacter('-', ' '); - - return cleaned; -} - -inline std::vector getFeaturedModels() -{ - return { - { "stability/text-to-audio", - "Stable Audio Text to Audio", - "Generate music, sound effects, or soundscapes from a text prompt.", - "Stability AI", - { "Generation" } }, - { "stability/audio-to-audio", - "Stable Audio Audio to Audio", - "Create variations or transfer style using text and audio conditioning.", - "Stability AI", - { "Generation", "Effects" } }, - { "teamup-tech/text2midi-symbolic-music-generation", - "Text2Midi", - "Generate symbolic MIDI music from a text description.", - "Hugging Face", - { "Generation" } }, - { "teamup-tech/demucs-source-separation", - "Demucs", - "Split a music recording into drums, bass, vocals, and instrumental stems.", - "Hugging Face", - { "Source Separation" } }, - { "teamup-tech/solo-piano-audio-to-midi-transcription", - "High Resolution Piano Transcription", - "Convert solo piano audio into a corresponding MIDI performance.", - "Hugging Face", - { "Analysis" } }, - { "teamup-tech/transkun", - "Transkun", - "Transcribe musical audio into symbolic note events.", - "Hugging Face", - { "Analysis" } }, - { "teamup-tech/TRIA", - "TRIA", - "Generate drum accompaniment conditioned on rhythmic input.", - "Hugging Face", - { "Performance Rendering and Synthesis", "Generation" } }, - { "teamup-tech/anticipatory-music-transformer", - "Anticipatory Music Transformer", - "Harmonize MIDI melodies by generating musically compatible notes.", - "Hugging Face", - { "Generation" } }, - { "teamup-tech/vampnet-conditional-music-generation", - "VampNet", - "Generate controllable variations of an input music recording.", - "Hugging Face", - { "Generation", "Effects" } }, - { "teamup-tech/harmonic-percussive-separation", - "Harmonic/Percussive Separation", - "Separate audio into harmonic and percussive components.", - "Hugging Face", - { "Source Separation" } }, - { "teamup-tech/Kokoro-TTS", - "Kokoro TTS", - "Generate speech from text using a selected voice preset.", - "Hugging Face", - { "Performance Rendering and Synthesis" } }, - { "teamup-tech/MegaTTS3-Voice-Cloning", - "MegaTTS3 Voice Cloning", - "Generate speech from text conditioned on a reference voice recording.", - "Hugging Face", - { "Performance Rendering and Synthesis" } }, - { "teamup-tech/midi-synthesizer", - "MIDI Synthesizer", - "Render MIDI into audio using the standard MuseScore SoundFont.", - "Hugging Face", - { "Performance Rendering and Synthesis" } }, - { "teamup-tech/audioseal", - "AudioSeal", - "Apply or inspect audio watermarking for generated audio workflows.", - "Hugging Face", - { "Analysis", "Production" } }, - }; -} - -inline std::vector getFeaturedModelPaths() -{ - std::vector paths { "click here to enter a custom path..." }; - - for (const auto& entry : getFeaturedModels()) - paths.push_back(entry.path.toStdString()); - - return paths; -} - -inline Entry getEntryForPath(const String& modelPath) -{ - const auto cleanedPath = getCleanModelPath(modelPath).upToFirstOccurrenceOf(" [", false, false).trim(); - - for (const auto& entry : getFeaturedModels()) - { - if (entry.path == cleanedPath) - { - auto result = entry; - result.path = cleanedPath; - return result; - } - } - - return { cleanedPath, - getFallbackModelDisplayName(modelPath), - "Custom or recently used HARP-compatible model endpoint.", - cleanedPath.startsWith("stability/") ? "Stability AI" : "Custom" }; -} -} // namespace ModelRegistry diff --git a/src/utils/ModelTags.h b/src/utils/ModelTags.h new file mode 100644 index 00000000..50ff839a --- /dev/null +++ b/src/utils/ModelTags.h @@ -0,0 +1,319 @@ +/** + * @file ModelTags.h + * @brief Model taxonomy and parsing of the tags models declare about themselves. + * + * Tags are plain strings, read from a Space's README metadata when browsing models and + * from the model card once a model is loaded. Structured tags take the form + * ":" (e.g., "category:separation"), and anything else is a custom tag. + * + * The taxonomy is pyharp's (pyharp/pyharp/taxonomy.json), which is where models declare + * their tags, embedded at build time. A tag it does not recognize is kept as a custom tag, + * so a model tagged against a newer taxonomy is still shown, just not categorized by it. + */ + +#pragma once + +#include +#include +#include + +#include + +#include + +using namespace juce; + +namespace Taxonomy +{ +struct Subcategory +{ + String id; + String displayName; +}; + +struct Category +{ + String id; + String displayName; + std::vector subcategories; +}; + +// In display order, as defined by pyharp's taxonomy.json (embedded when HARP is built) +inline const std::vector& getCategories() +{ + static const std::vector categories = [] + { + std::vector parsed; + + const var taxonomy = JSON::parse( + String::fromUTF8(TaxonomyData::taxonomy_json, TaxonomyData::taxonomy_jsonSize)); + + if (auto* categoryList = taxonomy["categories"].getArray()) + { + for (const auto& categoryData : *categoryList) + { + Category category { categoryData["id"].toString(), + categoryData["name"].toString(), + {} }; + + if (auto* subcategoryList = categoryData["subcategories"].getArray()) + { + for (const auto& subcategoryData : *subcategoryList) + { + category.subcategories.push_back( + { subcategoryData["id"].toString(), + subcategoryData["name"].toString() }); + } + } + + parsed.push_back(std::move(category)); + } + } + + // Malformed taxonomy data is a build error, which cannot be recovered at run time + jassert(! parsed.empty()); + + return parsed; + }(); + + return categories; +} + +inline const Category* findCategory(const String& id) +{ + for (const auto& category : getCategories()) + { + if (category.id == id) + return &category; + } + + return nullptr; +} + +// Subcategory ids are unique across the taxonomy, so each one names its category too +inline const Category* findCategoryOfSubcategory(const String& id, + const Subcategory** subcategory = nullptr) +{ + for (const auto& category : getCategories()) + { + for (const auto& candidate : category.subcategories) + { + if (candidate.id == id) + { + if (subcategory != nullptr) + *subcategory = &candidate; + + return &category; + } + } + } + + return nullptr; +} +} // namespace Taxonomy + +/** + * What a model's tags say about it. + */ +struct ModelTags +{ + // One kind of data a model takes in or gives back, from tags like "output:file/json" + struct DataType + { + String modality; // e.g. "audio", "midi", "text", "file", "labels" + StringArray formats; // File extensions it is restricted to, if any, e.g. "json" + + // e.g. "File (JSON, TXT)" + String describe() const + { + String name = modality == "midi" + ? String("MIDI") + : modality.substring(0, 1).toUpperCase() + modality.substring(1); + + if (! formats.isEmpty()) + name << " (" << formats.joinIntoString(", ").toUpperCase() << ")"; + + return name; + } + }; + + StringArray categories; // Category ids, in taxonomy order + StringArray subcategories; // Subcategory ids, in taxonomy order + std::vector inputs; + std::vector outputs; + int sampleRate = 0; // Hz, or 0 if not declared + String channels; // e.g. "mono", "stereo", "6" + StringArray custom; // Everything else, as given + + static ModelTags parse(const StringArray& tags) + { + ModelTags parsed; + + StringArray categoriesFound; + StringArray subcategoriesFound; + + for (auto tag : tags) + { + tag = tag.trim(); + + if (tag.isEmpty()) + continue; + + const String key = tag.upToFirstOccurrenceOf(":", false, false).toLowerCase(); + const String value = tag.fromFirstOccurrenceOf(":", false, false).trim(); + const bool hasValue = tag.containsChar(':') && value.isNotEmpty(); + + if (hasValue && key == "category" && Taxonomy::findCategory(value) != nullptr) + { + categoriesFound.addIfNotAlreadyThere(value); + } + else if (hasValue && key == "subcategory" + && Taxonomy::findCategoryOfSubcategory(value) != nullptr) + { + subcategoriesFound.addIfNotAlreadyThere(value); + // A subcategory implies its category + categoriesFound.addIfNotAlreadyThere( + Taxonomy::findCategoryOfSubcategory(value)->id); + } + else if (hasValue && (key == "input" || key == "output")) + { + addDataType(key == "input" ? parsed.inputs : parsed.outputs, value); + } + else if (hasValue && key == "sample-rate" && value.containsOnly("0123456789")) + { + parsed.sampleRate = value.getIntValue(); + } + else if (hasValue && key == "channels") + { + parsed.channels = value.toLowerCase(); + } + else + { + parsed.custom.addIfNotAlreadyThere(tag); + } + } + + // Keep taxonomy order, regardless of the order in which the tags were declared + for (const auto& category : Taxonomy::getCategories()) + { + if (categoriesFound.contains(category.id)) + parsed.categories.add(category.id); + + for (const auto& subcategory : category.subcategories) + { + if (subcategoriesFound.contains(subcategory.id)) + parsed.subcategories.add(subcategory.id); + } + } + + return parsed; + } + + bool isCategorized() const { return ! categories.isEmpty(); } + + bool isInCategory(const String& categoryId) const + { + return categories.contains(categoryId); + } + + // e.g. "Audio, File (JSON, TXT)", or empty when none are declared + static String describe(const std::vector& dataTypes) + { + StringArray descriptions; + + for (const auto& dataType : dataTypes) + descriptions.add(dataType.describe()); + + return descriptions.joinIntoString(", "); + } + + // Short labels for the most distinguishing tags, most important first + StringArray getDisplayLabels() const + { + StringArray labels; + + // The most specific place in the taxonomy: each subcategory, or the category itself + // where none of its subcategories is given (e.g., "Utility", which has none) + for (const auto& category : Taxonomy::getCategories()) + { + if (! categories.contains(category.id)) + continue; + + bool hasSubcategory = false; + + for (const auto& subcategory : category.subcategories) + { + if (! subcategories.contains(subcategory.id)) + continue; + + hasSubcategory = true; + + // "Music" alone is ambiguous, so the analysis subcategories carry their + // category + labels.add(category.id == "analysis" + ? subcategory.displayName + " " + category.displayName + : subcategory.displayName); + } + + if (! hasSubcategory) + labels.add(category.displayName); + } + + if (sampleRate > 0) + labels.add(sampleRate % 1000 == 0 ? String(sampleRate / 1000) + " kHz" + : String(sampleRate / 1000.0, 1) + " kHz"); + + if (channels.isNotEmpty()) + labels.add(channels.containsOnly("0123456789") + ? channels + " channels" + : channels.substring(0, 1).toUpperCase() + channels.substring(1)); + + labels.addArray(custom); + + return labels; + } + + // Everything a search should match against, lowercased + String getSearchableText() const + { + StringArray words; + + for (const auto& id : categories) + words.add(Taxonomy::findCategory(id)->displayName); + + words.addArray(getDisplayLabels()); + words.add(describe(inputs)); + words.add(describe(outputs)); + + return words.joinIntoString(" ").toLowerCase(); + } + +private: + /* A value is a modality, optionally followed by the formats it is restricted to, e.g. + "file/json|txt". Tags for the same kind of data (e.g., two file outputs) are merged. */ + static void addDataType(std::vector& dataTypes, const String& value) + { + const String modality = value.upToFirstOccurrenceOf("/", false, false).trim().toLowerCase(); + const StringArray formats = StringArray::fromTokens( + value.fromFirstOccurrenceOf("/", false, false).toLowerCase(), "|", ""); + + if (modality.isEmpty()) + return; + + auto existing = std::find_if(dataTypes.begin(), + dataTypes.end(), + [&](const DataType& d) { return d.modality == modality; }); + + if (existing == dataTypes.end()) + { + dataTypes.push_back({ modality, {} }); + existing = std::prev(dataTypes.end()); + } + + for (const auto& format : formats) + { + if (format.trim().isNotEmpty()) + existing->formats.addIfNotAlreadyThere(format.trim()); + } + } +}; diff --git a/src/utils/Tutorial.h b/src/utils/Tutorial.h deleted file mode 100644 index 835a0593..00000000 --- a/src/utils/Tutorial.h +++ /dev/null @@ -1,12 +0,0 @@ -/** - * @file Tutorial.h - * @brief Constants needed for tutorial steps. - * @author saumya-pailwan - */ - -#pragma once - -namespace TutorialConstants -{ -inline constexpr const char* fallbackModelPath = "teamup-tech/demucs-source-separation"; -} diff --git a/src/widgets/MediaClipboardWidget.h b/src/widgets/MediaClipboardWidget.h index 8526b63d..3fa60334 100644 --- a/src/widgets/MediaClipboardWidget.h +++ b/src/widgets/MediaClipboardWidget.h @@ -14,6 +14,8 @@ #include "../utils/Logging.h" +#include "../windows/tutorial/TutorialTargets.h" + using namespace juce; class MediaClipboardWidget : public Component, public ChangeListener @@ -28,6 +30,7 @@ class MediaClipboardWidget : public Component, public ChangeListener initializeButtons(); controlsComponent.addAndMakeVisible(buttonsComponent); addAndMakeVisible(controlsComponent); + controlsComponent.setComponentID(TutorialTargets::clipboardControls); resetState(); @@ -133,43 +136,6 @@ class MediaClipboardWidget : public Component, public ChangeListener trackAreaWidget.addTrackFromFilePath(filePath, fromDAW); } - Rectangle getClipboardTrackAreaBounds() const { return trackArea.getBounds().expanded(2); } - - Rectangle getClipboardControlsBounds() const - { - return controlsComponent.getBounds().expanded(2); - } - - Rectangle getClipboardNameBoxBounds() const - { - return getLocalArea(&controlsComponent, selectionTextBox.getBounds()).expanded(2); - } - - Rectangle getClipboardButtonsBounds() const - { - return getLocalArea(&controlsComponent, buttonsComponent.getBounds()).expanded(2); - } - - Rectangle getAddFileButtonBounds() const - { - return getLocalArea(&buttonsComponent, addFileButton.getBounds()).expanded(2); - } - - Rectangle getRemoveButtonBounds() const - { - return getLocalArea(&buttonsComponent, removeSelectionButton.getBounds()).expanded(2); - } - - Rectangle getPlayButtonBounds() const - { - return getLocalArea(&buttonsComponent, playStopButton.getBounds()).expanded(2); - } - - Rectangle getSendToDAWButtonBounds() const - { - return getLocalArea(&buttonsComponent, sendToDAWButton.getBounds()).expanded(2); - } - void addFileCallback() { StringArray validExtensions = MediaDisplayComponent::getSupportedExtensions(); diff --git a/src/widgets/ModelInfoWidget.h b/src/widgets/ModelInfoWidget.h index 5ea75411..23154668 100644 --- a/src/widgets/ModelInfoWidget.h +++ b/src/widgets/ModelInfoWidget.h @@ -6,9 +6,11 @@ #pragma once -#include #include +#include + +#include "../widgets/ModelStyle.h" #include "../widgets/StatusAreaWidget.h" #include "../gui/HoverHandler.h" @@ -21,29 +23,21 @@ using namespace juce; class ModelAuthorLabel : public Component { public: - ModelAuthorLabel(const String& modelName = "", - const String& author = "", - const URL& newURL = URL()) + ModelAuthorLabel() { - if (modelName.isNotEmpty()) - { - setModelName(modelName); - } - - if (author.isNotEmpty()) - { - setAuthor(author); - } - - setURL(newURL); - - modelLabel.setFont(Font(22.0f, Font::bold)); + modelLabel.setFont(ModelStyle::font(18.0f, true)); + modelLabel.setBorderSize({ 0, 0, 0, 0 }); modelLabel.onHover = [this] { instructionsMessage->setMessage("Click to view the model's webpage or documentation."); }; modelLabel.onExit = [this] { instructionsMessage->clearMessage(); }; modelLabel.onClick = [this] { url.launchInDefaultBrowser(); }; + authorLabel.setFont(ModelStyle::font(13.0f)); + authorLabel.setColour(Label::textColourId, ModelStyle::secondaryText); + + setURL(URL()); + addAndMakeVisible(modelLabel); addAndMakeVisible(authorLabel); } @@ -52,10 +46,11 @@ class ModelAuthorLabel : public Component { Rectangle totalArea = getLocalBounds(); - float modelNameWidth = modelLabel.getFont().getStringWidthFloat(modelLabel.getText()) + 10; + const int modelNameWidth = + ModelStyle::getTextWidth(modelLabel.getFont(), modelLabel.getText()) + 2; - modelLabel.setBounds(totalArea.removeFromLeft(static_cast(modelNameWidth))); - authorLabel.setBounds(totalArea.translated(0, 3)); + modelLabel.setBounds(totalArea.removeFromLeft(jmin(modelNameWidth, totalArea.getWidth()))); + authorLabel.setBounds(totalArea.withTrimmedTop(2)); } void setModelName(const String& modelName) @@ -64,28 +59,22 @@ class ModelAuthorLabel : public Component resized(); } + void setAuthor(const String& author) { authorLabel.setText(author, dontSendNotification); resized(); } + void setURL(const URL& newURL) { - if (newURL.isWellFormed()) - { - url = newURL; + const bool isLinked = newURL.isWellFormed() && newURL.toString(false).isNotEmpty(); - modelLabel.setHoverColor(Colours::coral); - modelLabel.setHoverable(true); + url = newURL; - resized(); - } - else - { - modelLabel.setHoverColor(Colours::white); - modelLabel.setHoverable(false); - } + modelLabel.setHoverColor(isLinked ? ModelStyle::accent : Colours::white); + modelLabel.setHoverable(isLinked); } private: @@ -97,6 +86,10 @@ class ModelAuthorLabel : public Component SharedResourcePointer instructionsMessage; }; +/** + * The loaded model's card, presented like the Home tab's: its name (linking to its page) + * and author, notes on its deployment, its description, and its tags. + */ class ModelInfoWidget : public Component { public: @@ -104,146 +97,126 @@ class ModelInfoWidget : public Component { addAndMakeVisible(modelAuthorLabel); - // Configure description as scrollable read-only text + // Scrollable read-only text, for descriptions longer than the widget shows description.setMultiLine(true); description.setReadOnly(true); description.setScrollbarsShown(true); - description.setCaretVisible(false); // no blinking cursor - description.setPopupMenuEnabled(false); // disable right-click menu - description.setFont(Font(15.0f)); - - // Make it visually match the old TextLabel appearance + description.setCaretVisible(false); + description.setPopupMenuEnabled(false); + description.setFont(descriptionFont); + description.setIndents(0, 0); + description.setColour(TextEditor::textColourId, Colours::whitesmoke.withAlpha(0.85f)); description.setColour(TextEditor::backgroundColourId, Colours::transparentBlack); description.setColour(TextEditor::outlineColourId, Colours::transparentBlack); + description.setColour(TextEditor::focusedOutlineColourId, Colours::transparentBlack); description.setColour(TextEditor::shadowColourId, Colours::transparentBlack); - addAndMakeVisible(description); - } - ~ModelInfoWidget() {} + addChildComponent(tagRow); + } - //void paint(Graphics& g) {} + void paint(Graphics& g) override + { + ModelStyle::drawCard(g, getLocalBounds().toFloat().reduced(1.0f)); + ModelStyle::drawBadges(g, badges, getContentArea().removeFromTop(headerHeight)); + } void resized() override { - Rectangle bounds = getLocalBounds().reduced(marginSize); + auto area = getContentArea(); + + auto headerRow = area.removeFromTop(headerHeight); + headerRow.setRight(ModelStyle::getBadgesLeft(badges, headerRow) - ModelStyle::chipGap); + modelAuthorLabel.setBounds(headerRow); - // Set fixed size for model and author labels - modelAuthorLabel.setBounds(bounds.removeFromTop(headerHeight)); - // Grant remaining space to description - description.setBounds(bounds); + area.removeFromTop(rowGap); + description.setBounds(area.removeFromTop(getDescriptionHeightForWidth(area.getWidth()))); + + if (! tagRow.isEmpty()) + { + area.removeFromTop(rowGap + 2); + tagRow.setBounds(area.removeFromTop(ModelStyle::chipHeight)); + } } int getPreferredHeightForWidth(int width) const { - const int contentWidth = jmax(120, width - 2 * (int) marginSize); - const int lineCount = estimateWrappedLineCount(description.getText(), contentWidth); - const int visibleLines = jlimit(1, 4, lineCount); - const int lineHeight = (int) std::ceil(description.getFont().getHeight() + 2.0f); - const int descriptionHeight = visibleLines * lineHeight + 4; + const int contentWidth = width - 2 * horizontalPadding; - return (int) (2 * marginSize + headerHeight + descriptionHeight); - } + int height = 2 * verticalPadding + headerHeight + rowGap + + getDescriptionHeightForWidth(contentWidth); - void resetState() - { - ModelMetadata emptyMetadata; + if (! tagRow.isEmpty()) + height += rowGap + 2 + ModelStyle::chipHeight; - updateLabels(emptyMetadata); - modelAuthorLabel.setURL(URL("")); + return height; } void updateLabels(const ModelMetadata& metadata) { - if (metadata.name.empty()) - { - modelAuthorLabel.setModelName(""); - } - else - { - modelAuthorLabel.setModelName(String(metadata.name)); - } + modelAuthorLabel.setModelName(String(metadata.name)); + modelAuthorLabel.setAuthor(metadata.author.empty() ? String() + : "by " + String(metadata.author)); - if (metadata.author.empty()) - { - modelAuthorLabel.setAuthor(""); - } - else - { - modelAuthorLabel.setAuthor("by " + String(metadata.author)); - } + description.setText(String(metadata.description)); - if (metadata.description.empty()) - { - description.setText(""); - } - else - { - description.setText(String(metadata.description)); - } + StringArray tags; + + for (const auto& tag : metadata.tags) + tags.add(tag); + + tagRow.setTags(ModelTags::parse(tags)); + tagRow.setVisible(! tagRow.isEmpty()); resized(); + repaint(); + } + + void setBadges(std::vector newBadges) + { + badges = std::move(newBadges); + + resized(); + repaint(); } void addOpenablePath(const String& openablePath) { modelAuthorLabel.setURL(URL(openablePath)); } private: - int estimateWrappedLineCount(const String& text, int availableWidth) const + Rectangle getContentArea() const + { + return getLocalBounds().reduced(horizontalPadding, verticalPadding); + } + + // Tall enough for the description, up to a limit beyond which it scrolls + int getDescriptionHeightForWidth(int width) const { - if (availableWidth <= 0) - return 1; + if (description.isEmpty() || width <= 0) + return 0; - auto font = description.getFont(); - const float spaceWidth = jmax(1.0f, font.getStringWidthFloat(" ")); - int lines = 0; + AttributedString text; + text.append(description.getText(), descriptionFont); - StringArray paragraphs; - paragraphs.addLines(text.isEmpty() ? String(" ") : text); + TextLayout layout; + // Leave room for the scrollbar, which appears once the text is too long + layout.createLayout(text, (float) (width - getLookAndFeel().getDefaultScrollbarWidth())); - for (const auto& paragraph : paragraphs) - { - StringArray words; - words.addTokens(paragraph, " ", ""); - words.removeEmptyStrings(); - - if (words.isEmpty()) - { - ++lines; - continue; - } - - float currentLineWidth = 0.0f; - for (const auto& word : words) - { - const float wordWidth = font.getStringWidthFloat(word); - - if (currentLineWidth <= 0.0f) - { - currentLineWidth = wordWidth; - continue; - } - - if (currentLineWidth + spaceWidth + wordWidth <= (float) availableWidth) - { - currentLineWidth += spaceWidth + wordWidth; - } - else - { - ++lines; - currentLineWidth = wordWidth; - } - } - - ++lines; - } + const int lines = jlimit(1, maxDescriptionLines, layout.getNumLines()); - return jmax(1, lines); + return (int) std::ceil((float) lines * descriptionFont.getHeight()) + 2; } - const float headerHeight = 30; - const float marginSize = 2; + static constexpr int horizontalPadding = 12; + static constexpr int verticalPadding = 9; + static constexpr int headerHeight = 24; + static constexpr int rowGap = 2; + static constexpr int maxDescriptionLines = 4; - ModelAuthorLabel modelAuthorLabel; + const Font descriptionFont = ModelStyle::font(13.0f); + ModelAuthorLabel modelAuthorLabel; TextEditor description; + ModelStyle::TagRow tagRow; + + std::vector badges; }; diff --git a/src/widgets/ModelSelectionWidget.h b/src/widgets/ModelSelectionWidget.h deleted file mode 100644 index 09eb4b84..00000000 --- a/src/widgets/ModelSelectionWidget.h +++ /dev/null @@ -1,654 +0,0 @@ -/** - * @file ModelSelectionWidget.h - * @brief Component allowing for selection and loading of model. - * @author hugofloresgarcia, rc2000123, xribene, lindseydeng, cwitkowitz - */ - -#pragma once - -#include -#include -#include - -#include - -#include "../widgets/StatusAreaWidget.h" - -#include "../gui/HoverHandler.h" -#include "../gui/MultiButton.h" - -#include "../utils/Clients.h" -#include "../utils/Errors.h" -#include "../utils/Interface.h" -#include "../utils/Logging.h" -#include "../utils/ModelRegistry.h" - -using namespace juce; - -struct SharedChoices : public ChangeBroadcaster -{ - enum class LoadStatus - { - None, - Error, - Down, - TryAgain - }; - - SharedChoices() - : savedModelPaths(ModelRegistry::getFeaturedModelPaths()) - { - } - - int getIndexForPath(const std::string& p) - { - int idx = -1; - const auto cleanedPath = stripStatusTag(p); - - for (unsigned int i = 0; i < savedModelPaths.size(); ++i) - { - if (savedModelPaths[i] == cleanedPath) - { - idx = (int) i; - - break; - } - } - - return idx; - } - - // TODO - should check endpoint path - otherwise could be duplicates - bool containsPath(const std::string& p) { return getIndexForPath(p) != -1; } - - void addNewPath(const std::string& p) - { - const auto cleanedPath = stripStatusTag(p); - - if (! containsPath(cleanedPath)) - savedModelPaths.push_back(cleanedPath); - - sendSynchronousChangeMessage(); - } - - void updatePath(unsigned int idx, const std::string& p) - { - savedModelPaths[idx] = stripStatusTag(p); - sendSynchronousChangeMessage(); - } - - void setLoadStatus(const std::string& p, LoadStatus status) - { - const auto cleanedPath = stripStatusTag(p); - - if (! containsPath(cleanedPath)) - savedModelPaths.push_back(cleanedPath); - - if (status == LoadStatus::None) - loadStatuses.erase(cleanedPath); - else - loadStatuses[cleanedPath] = status; - - sendSynchronousChangeMessage(); - } - - String getDisplayTextForIndex(unsigned int idx) const - { - if (idx >= savedModelPaths.size()) - return {}; - - const auto& path = savedModelPaths[idx]; - const auto status = loadStatuses.find(path); - - if (status == loadStatuses.end()) - return path; - - return String(path) + getStatusTag(status->second); - } - - static std::string stripStatusTag(const std::string& p) - { - auto cleaned = String(p).trim(); - - for (const auto& tag : { errorTag, downTag, tryAgainTag }) - cleaned = cleaned.replace(String(tag), ""); - - return cleaned.trim().toStdString(); - } - - std::vector savedModelPaths; - -private: - static String getStatusTag(LoadStatus status) - { - switch (status) - { - case LoadStatus::Error: - return errorTag; - case LoadStatus::Down: - return downTag; - case LoadStatus::TryAgain: - return tryAgainTag; - case LoadStatus::None: - default: - return {}; - } - } - - inline static const String errorTag { " [ERROR]" }; - inline static const String downTag { " [DOWN]" }; - inline static const String tryAgainTag { " [TRY AGAIN]" }; - - std::map loadStatuses; -}; - -class CustomPathComponent : public Component -{ -public: - CustomPathComponent(std::function onLoad, std::function onCancel) - : onLoadCallback(std::move(onLoad)), onCancelCallback(std::move(onCancel)) - { - pathEditor.setMultiLine(false); - pathEditor.setReturnKeyStartsNewLine(false); - pathEditor.onTextChange = [this] { loadButton.setEnabled(! pathEditor.isEmpty()); }; - pathEditor.onReturnKey = [this] - { - if (loadButton.isEnabled()) - { - loadButton.triggerClick(); - } - }; - addAndMakeVisible(pathEditor); - - loadButton.setEnabled(false); - loadButton.onClick = [this] - { - wasLoadPressed = true; - - if (onLoadCallback) - { - onLoadCallback(pathEditor.getText()); - } - - closePopup(); - }; - addAndMakeVisible(loadButton); - - cancelButton.onClick = [this] { closePopup(); }; - addAndMakeVisible(cancelButton); - - setSize(popupWidth, popupHeight); - } - - ~CustomPathComponent() override - { - if (! wasLoadPressed && onCancelCallback) - { - // Treat as cancel if closed without load - onCancelCallback(); - } - } - - void visibilityChanged() override - { - if (isVisible()) - { - MessageManager::callAsync([this] { pathEditor.grabKeyboardFocus(); }); - } - } - - void resized() override - { - Rectangle fullArea = getLocalBounds(); - - /* TopLevelWindow::centreAroundComponent shrinks a dialog to fit the area it - is centered within, so this can be laid out smaller than the size asked - for. Fixed items would then exceed the space and leave the rest negative. */ - if (fullArea.isEmpty()) - { - pathEditor.setBounds({}); - loadButton.setBounds({}); - cancelButton.setBounds({}); - - return; - } - - const int editorHeight = jmin(editorRowHeight, fullArea.getHeight()); - - FlexBox fullPopup; - fullPopup.flexDirection = FlexBox::Direction::column; - - fullPopup.items.add(FlexItem(pathEditor) - .withHeight((float) editorHeight) - .withMargin(jmin(2.0f, (float) fullArea.getHeight() / 4.0f))); - - FlexBox buttonsArea; - buttonsArea.flexDirection = FlexBox::Direction::row; - - // Margins have to shrink with the popup, or they consume more than there is - const float buttonMargin = - jmin(10.0f, (float) jmin(fullArea.getWidth(), fullArea.getHeight()) / 8.0f); - - buttonsArea.items.add(FlexItem(loadButton).withFlex(1).withMargin(buttonMargin)); - buttonsArea.items.add(FlexItem().withFlex(0.25)); - buttonsArea.items.add(FlexItem(cancelButton).withFlex(1).withMargin(buttonMargin)); - - fullPopup.items.add(FlexItem(buttonsArea).withFlex(1).withMinHeight(0.0f)); - - fullPopup.performLayout(fullArea); - - // FlexBox can still hand out negative sizes when the space runs out - for (auto* child : getChildren()) - { - if (child->getWidth() < 0 || child->getHeight() < 0) - { - child->setBounds(child->getBounds().withSize(jmax(0, child->getWidth()), - jmax(0, child->getHeight()))); - } - } - } - - void paint(Graphics& g) override - { - g.fillAll(getUIColourIfAvailable(LookAndFeel_V4::ColourScheme::UIColour::windowBackground)); - } - - void setTextFieldValue(const String& text) - { - pathEditor.setText(text, dontSendNotification); - pathEditor.selectAll(); - } - -private: - void closePopup() - { - if (auto* popup = findParentComponentOfClass()) - { - popup->exitModalState(0); - } - } - - static constexpr int popupWidth = 400; - static constexpr int popupHeight = 80; - static constexpr int editorRowHeight = 30; - - TextEditor pathEditor; - TextButton loadButton { "Load" }; - TextButton cancelButton { "Cancel" }; - - bool wasLoadPressed = false; - - std::function onLoadCallback; - std::function onCancelCallback; -}; - -class ModelSelectionWidget : public Component, public ChangeBroadcaster, public ChangeListener -{ -public: - ModelSelectionWidget() - { - initializeLoadModelButton(); - initializeModelPathComboBox(); - - resetState(); - - sharedChoices->addChangeListener(this); - } - - ~ModelSelectionWidget() { sharedChoices->removeChangeListener(this); } - - void resized() override - { - FlexBox selectionArea; - selectionArea.flexDirection = FlexBox::Direction::row; - - selectionArea.items.add(FlexItem(modelPathComboBox).withFlex(1).withMargin(marginSize)); - selectionArea.items.add(FlexItem(loadModelButton).withWidth(100).withMargin(marginSize)); - - selectionArea.performLayout(getLocalBounds()); - } - - String getCurrentlySelectedPath() { return selectedPath; } - - void loadModelBypass(const String& modelPath) - { - /* Reduce to the canonical form on the way in, so that the same model - entered three different ways is loaded, listed, and matched as one - entry rather than accumulating duplicates in the dropdown. */ - selectedPath = canonicalizeModelPath(SharedChoices::stripStatusTag(modelPath.toStdString())); - - if (selectedPath != modelPath) - { - DBG_AND_LOG("ModelSelectionWidget::loadModelBypass: Path \"" - << modelPath << "\" resolved to \"" << selectedPath << "\"."); - } - - sendChangeMessage(); - } - - void resetState() - { - lastLoadedPathIndex = -1; - lastSelectedPathIndex = -1; - modelPathComboBox.setSelectedId(lastSelectedPathIndex); - modelPathComboBox.setEnabled(true); - - loadModelButton.setEnabled(false); - } - - void setDisabled() - { - modelPathComboBox.setEnabled(false); - loadModelButton.setEnabled(false); - } - - void setEnabled() - { - modelPathComboBox.setEnabled(true); - loadModelButton.setEnabled(true); - } - - void setFinishedState() - { - setEnabled(); - - modelPathComboBox.setSelectedId(lastLoadedPathIndex + 1); - } - - /* The path a model actually loaded from can differ from the one entered, since - the provider is asked for its exact spelling. The dropdown lists that one, so - the three ways of writing an address collapse to a single entry. */ - void setSuccessfulState(const String& resolvedPath) - { - if (resolvedPath.isNotEmpty()) - { - selectedPath = resolvedPath; - } - - std::string loadedPath = SharedChoices::stripStatusTag(selectedPath.toStdString()); - - if (! sharedChoices->containsPath(loadedPath)) - { - // Add a new entry for custom path - sharedChoices->addNewPath(loadedPath); - - lastSelectedPathIndex = sharedChoices->getIndexForPath(loadedPath); - } - - sharedChoices->setLoadStatus(loadedPath, SharedChoices::LoadStatus::None); - - lastLoadedPathIndex = sharedChoices->getIndexForPath(loadedPath); - - setFinishedState(); - } - - void setUnsuccessfulState(const Error& error) - { - bool wasValidPath = true; - - if (const auto* e = std::get_if(&error)) - { - wasValidPath = false; - } - - if (const auto* e = std::get_if(&error)) - { - if (e->type == HttpError::Type::BadStatusCode && e->statusCode == 404) - { - wasValidPath = false; - } - } - - if (! wasValidPath) - { - if (modelPathComboBox.getSelectedItemIndex() == 0) - { - openCustomPathPopup(selectedPath); - - return; - } - } - - std::string originalEntry = SharedChoices::stripStatusTag(selectedPath.toStdString()); - auto status = SharedChoices::LoadStatus::Error; - - if (const auto* e = std::get_if(&error)) - { - if (e->type == HttpError::Type::ConnectionFailed - && e->request == HttpError::Request::POST) - { - status = SharedChoices::LoadStatus::TryAgain; - } - else if (e->type == HttpError::Type::BadStatusCode && e->statusCode == 429) - { - // Rate limiting says nothing about the model, only about how - // quickly it was asked for - status = SharedChoices::LoadStatus::TryAgain; - } - else if (e->type == HttpError::Type::BadStatusCode && e->statusCode == 503) - { - status = SharedChoices::LoadStatus::Down; - } - } - else if (const auto* e = std::get_if(&error)) - { - /* A Space that is waking up is not a broken one, so it must not be - marked as down. The Hub was queried to determine the stage. */ - if (e->type == GradioError::Type::SpaceStarting) - { - status = SharedChoices::LoadStatus::TryAgain; - } - else if (e->type == GradioError::Type::SpaceUnavailable) - { - status = SharedChoices::LoadStatus::Down; - } - } - - sharedChoices->setLoadStatus(originalEntry, status); - - lastSelectedPathIndex = lastLoadedPathIndex; - - selectedPath.clear(); - - setFinishedState(); - } - - void changeListenerCallback(ChangeBroadcaster* /*source*/) { resetModelPathComboBox(); } - -private: - void resetModelPathComboBox() - { - modelPathComboBox.clear(); - - for (unsigned int i = 0; i < sharedChoices->savedModelPaths.size(); ++i) - { - // Add saved path to combo box (skipping 0 for custom path) - modelPathComboBox.addItem(sharedChoices->getDisplayTextForIndex(i), - static_cast(i) + 1); - } - } - - void initializeModelPathComboBox() - { - modelPathComboBox.setTextWhenNothingSelected("click here to select a model..."); - - resetModelPathComboBox(); - - modelPathComboBox.onChange = [this] - { - if (modelPathComboBox.getSelectedItemIndex() == -1) - { - DBG_AND_LOG("ModelSelectionWidget::modelPathComboBox::onChange: Combo box reset."); - } - else - { - if (modelPathComboBox.getSelectedItemIndex() == 0) - { - DBG_AND_LOG( - "ModelSelectionWidget::modelPathComboBox::onChange: Custom path selected."); - - openCustomPathPopup(); - } - else - { - lastSelectedPathIndex = modelPathComboBox.getSelectedItemIndex(); - - DBG_AND_LOG("ModelSelectionWidget::modelPathComboBox::onChange: Entry " - << lastSelectedPathIndex << " selected."); - } - - loadModelButton.setEnabled(true); - } - }; - - addAndMakeVisible(modelPathComboBox); - - modelPathComboBoxHandler.onMouseEnter = [this]() - { - if (instructionsMessage != nullptr) - { - instructionsMessage->setMessage( - "A drop-down menu with featured available models. Any custom paths " - "successfully loaded will automatically be added to the list."); - } - }; - modelPathComboBoxHandler.onMouseExit = [this]() - { - if (instructionsMessage != nullptr) - { - instructionsMessage->clearMessage(); - } - }; - modelPathComboBoxHandler.attach(); - } - - void initializeLoadModelButton() - { - std::function loadCallback = [this]() - { - if (modelPathComboBox.getSelectedItemIndex() != 0) - { - const auto selectedIndex = modelPathComboBox.getSelectedItemIndex(); - selectedPath = sharedChoices->savedModelPaths[(unsigned int) selectedIndex]; - - sendChangeMessage(); - } - }; - - // Mode when a model is selected and not currently being loaded (load enabled) - loadButtonActiveInfo = MultiButton::Mode { "Load", - "Click to load currently selected model path.", - loadCallback, - MultiButton::DrawingMode::TextOnly }; - loadModelButton.addMode(loadButtonActiveInfo); - loadModelButton.setMode(loadButtonActiveInfo.displayLabel); - addAndMakeVisible(loadModelButton); - } - - /** - * Create caollbacks for and launch the custom path popup. - */ - void openCustomPathPopup(const String& prefillText = "") - { - std::function loadCallback = [this](String path) - { - DBG_AND_LOG("ModelSelectionWidget::openCustomPathPopup::loadCallback: " - << "Custom path \"" << path << "\" entered."); - - loadModelBypass(path); - }; - - std::function cancelCallback = [this]() - { - DBG_AND_LOG("ModelSelectionWidget::openCustomPathPopup::cancelCallback: " - << "Custom path selection canceled."); - - if (lastLoadedPathIndex >= 0) - { - // Set combo box selection to last successfully loaded model - modelPathComboBox.setSelectedId(lastLoadedPathIndex + 1); - modelPathComboBox.setEnabled(true); - } - else - { - resetState(); - } - }; - - CustomPathComponent* content = - new CustomPathComponent(std::move(loadCallback), std::move(cancelCallback)); - - if (prefillText.isNotEmpty()) - { - content->setTextFieldValue(prefillText); - } - - DialogWindow::LaunchOptions options; - options.dialogTitle = "Enter Custom Path"; - options.dialogBackgroundColour = Colours::darkgrey; - options.content.setOwned(content); - - options.useNativeTitleBar = false; - options.resizable = false; - options.escapeKeyTriggersCloseButton = true; - options.componentToCentreAround = getParentComponent(); - - options.launchAsync(); - } - - const float marginSize = 2; - - SharedResourcePointer sharedChoices; - - /* A ComboBox that always opens its popup from the top of the list, showing all - items without scrolling to the currently-selected one first. Item 0 is the - custom path entry, which must stay reachable however far down the list the - loaded model sits. */ - struct FullListComboBox : public ComboBox - { - void showPopup() override - { - auto& lf = getLookAndFeel(); - auto label = std::unique_ptr