diff --git a/Sources/MusicAssistantApp/Presentation/FinalReviewView.swift b/Sources/MusicAssistantApp/Presentation/FinalReviewView.swift index ebc8d55..8a7e80c 100644 --- a/Sources/MusicAssistantApp/Presentation/FinalReviewView.swift +++ b/Sources/MusicAssistantApp/Presentation/FinalReviewView.swift @@ -156,9 +156,12 @@ struct FinalReviewView: View { private func ensureSunoOutput() { guard project.sunoOutput == nil else { return } + let compiledOutput = (try? SongProjectPromptCompiler().compile(project: project)) + ?? CompiledSunoOutput(lyricsText: project.lyrics.text, stylePrompt: "") + project.sunoOutput = SunoOutput( - lyricsText: project.lyrics.text, - stylePrompt: "", + lyricsText: compiledOutput.lyricsText, + stylePrompt: compiledOutput.stylePrompt, generatedAt: Date() ) } diff --git a/Sources/MusicAssistantApp/Presentation/InstrumentBrowserView.swift b/Sources/MusicAssistantApp/Presentation/InstrumentBrowserView.swift index c23c4af..73eb70b 100644 --- a/Sources/MusicAssistantApp/Presentation/InstrumentBrowserView.swift +++ b/Sources/MusicAssistantApp/Presentation/InstrumentBrowserView.swift @@ -3,18 +3,23 @@ import SwiftUI struct InstrumentBrowserView: View { @Binding private var project: SongProject + private let saveProject: (SongProject) async -> Bool private let catalog: LocalInstrumentCatalog @State private var searchText = "" @State private var selectedFamilyCategory: String? @State private var selectedRegionOrigin: String? + @State private var hasUnsavedSelectionChanges = false + @State private var isSaving = false @Environment(\.dismiss) private var dismiss init( project: Binding, + saveProject: @escaping (SongProject) async -> Bool = { _ in true }, catalog: LocalInstrumentCatalog = LocalInstrumentCatalog() ) { _project = project + self.saveProject = saveProject self.catalog = catalog } @@ -81,12 +86,33 @@ struct InstrumentBrowserView: View { .toolbar { ToolbarItem(placement: .cancellationAction) { Button("Done") { - dismiss() + Task { + await saveSelectionChangesIfNeeded() + dismiss() + } + } + .disabled(isSaving) + } + + if isSaving { + ToolbarItem(placement: .status) { + ProgressView() + .controlSize(.small) } } } } .frame(minWidth: 460, minHeight: 520) + .onDisappear { + guard hasUnsavedSelectionChanges, !isSaving else { return } + + let projectToSave = project + hasUnsavedSelectionChanges = false + + Task { + _ = await saveProject(projectToSave) + } + } } private func selectionBinding(for instrument: InstrumentCatalogItem) -> Binding { @@ -98,8 +124,18 @@ struct InstrumentBrowserView: View { isSelected: isSelected, variant: instrument.name ) + hasUnsavedSelectionChanges = true } } + + private func saveSelectionChangesIfNeeded() async { + guard hasUnsavedSelectionChanges else { return } + + isSaving = true + let didSave = await saveProject(project) + hasUnsavedSelectionChanges = !didSave + isSaving = false + } } #Preview { diff --git a/Sources/MusicAssistantApp/Presentation/ProjectInspectorView.swift b/Sources/MusicAssistantApp/Presentation/ProjectInspectorView.swift index 8e4374e..b9c058e 100644 --- a/Sources/MusicAssistantApp/Presentation/ProjectInspectorView.swift +++ b/Sources/MusicAssistantApp/Presentation/ProjectInspectorView.swift @@ -146,7 +146,7 @@ struct ProjectInspectorView: View { } .background(Color(nsColor: .windowBackgroundColor)) .sheet(isPresented: $isInstrumentBrowserPresented) { - InstrumentBrowserView(project: $project) + InstrumentBrowserView(project: $project, saveProject: saveProject) } } diff --git a/Sources/MusicAssistantCore/Services/AI/AIService.swift b/Sources/MusicAssistantCore/Services/AI/AIService.swift index 2f5cd49..89bea85 100644 --- a/Sources/MusicAssistantCore/Services/AI/AIService.swift +++ b/Sources/MusicAssistantCore/Services/AI/AIService.swift @@ -65,15 +65,20 @@ public struct SongProjectGenerationRequest: Equatable, Sendable { public var context: AIRequestContext public var seedProject: SongProject? public var discussionMode: ConversationMode + public var songGenerationContext: SongGenerationContext public init( context: AIRequestContext, seedProject: SongProject? = nil, - discussionMode: ConversationMode = .auto + discussionMode: ConversationMode = .auto, + songGenerationContext: SongGenerationContext? = nil ) { self.context = context self.seedProject = seedProject self.discussionMode = discussionMode + self.songGenerationContext = songGenerationContext + ?? seedProject.map(SongGenerationContext.init(project:)) + ?? .empty } } @@ -96,10 +101,16 @@ public struct SongProjectGenerationResult: Equatable, Sendable { public struct SongProjectDiscussionRequest: Equatable, Sendable { public var context: AIRequestContext public var project: SongProject + public var songGenerationContext: SongGenerationContext - public init(context: AIRequestContext, project: SongProject) { + public init( + context: AIRequestContext, + project: SongProject, + songGenerationContext: SongGenerationContext? = nil + ) { self.context = context self.project = project + self.songGenerationContext = songGenerationContext ?? SongGenerationContext(project: project) } } @@ -118,17 +129,20 @@ public struct LyricsRevisionRequest: Equatable, Sendable { public var project: SongProject public var sourceLyrics: Lyrics public var mode: LyricsRevisionMode + public var songGenerationContext: SongGenerationContext public init( context: AIRequestContext, project: SongProject, sourceLyrics: Lyrics, - mode: LyricsRevisionMode + mode: LyricsRevisionMode, + songGenerationContext: SongGenerationContext? = nil ) { self.context = context self.project = project self.sourceLyrics = sourceLyrics self.mode = mode + self.songGenerationContext = songGenerationContext ?? SongGenerationContext(project: project) } } @@ -154,15 +168,18 @@ public struct SongProjectUpdateRequest: Equatable, Sendable { public var context: AIRequestContext public var project: SongProject public var allowedScopes: [SongProjectUpdateScope] + public var songGenerationContext: SongGenerationContext public init( context: AIRequestContext, project: SongProject, - allowedScopes: [SongProjectUpdateScope] + allowedScopes: [SongProjectUpdateScope], + songGenerationContext: SongGenerationContext? = nil ) { self.context = context self.project = project self.allowedScopes = allowedScopes + self.songGenerationContext = songGenerationContext ?? SongGenerationContext(project: project) } } diff --git a/Sources/MusicAssistantCore/Services/PromptCompiler/PromptCompiling.swift b/Sources/MusicAssistantCore/Services/PromptCompiler/PromptCompiling.swift index 60cc6f7..1a805a2 100644 --- a/Sources/MusicAssistantCore/Services/PromptCompiler/PromptCompiling.swift +++ b/Sources/MusicAssistantCore/Services/PromptCompiler/PromptCompiling.swift @@ -13,3 +13,83 @@ public struct CompiledSunoOutput: Equatable, Sendable { self.stylePrompt = stylePrompt } } + +public struct SongProjectPromptCompiler: PromptCompiling { + public init() {} + + public func compile(project: SongProject) throws -> CompiledSunoOutput { + CompiledSunoOutput( + lyricsText: lyricsText(for: project), + stylePrompt: stylePrompt(for: project) + ) + } + + private func lyricsText(for project: SongProject) -> String { + if !project.lyrics.text.trimmingCharacters(in: .whitespacesAndNewlines).isEmpty { + return project.lyrics.text + } + + return project.orderedSections + .compactMap { section in + let lyrics = project.lyrics.sectionTexts[section.id] ?? section.lyrics + let trimmedLyrics = lyrics.trimmingCharacters(in: .whitespacesAndNewlines) + guard !trimmedLyrics.isEmpty else { return nil } + return "[\(section.title)]\n\(trimmedLyrics)" + } + .joined(separator: "\n\n") + } + + private func stylePrompt(for project: SongProject) -> String { + let context = SongGenerationContext(project: project) + var components: [String] = [] + + append(project.genres.map(\.name).joined(separator: " + "), to: &components) + append(project.moods.map(\.name).joined(separator: ", "), to: &components) + append(context.sunoStyleInstrumentPhrase, to: &components) + append(vocalPhrase(for: project), to: &components) + append(musicalParametersPhrase(for: project), to: &components) + append(project.productionDirections.map(\.text).joined(separator: ", "), to: &components) + + return components.joined(separator: "; ") + } + + private func vocalPhrase(for project: SongProject) -> String { + project.vocalists + .map { vocalist in + [ + vocalist.label, + vocalist.voiceType, + vocalist.genderSelection, + vocalist.performanceStyle + ] + .compactMap { normalized($0) } + .joined(separator: " ") + } + .filter { !$0.isEmpty } + .joined(separator: ", ") + } + + private func musicalParametersPhrase(for project: SongProject) -> String { + var parameters: [String] = [] + + if let bpm = project.bpm?.value { + parameters.append("\(bpm) BPM") + } + + append(project.key?.value.map { "key \($0)" }, to: ¶meters) + append(project.scale?.value.map { "\($0) scale" }, to: ¶meters) + append(project.maqam?.value.map { "maqam \($0)" }, to: ¶meters) + + return parameters.joined(separator: ", ") + } + + private func append(_ value: String?, to components: inout [String]) { + guard let normalizedValue = normalized(value) else { return } + components.append(normalizedValue) + } + + private func normalized(_ value: String?) -> String? { + let trimmedValue = value?.trimmingCharacters(in: .whitespacesAndNewlines) + return trimmedValue?.isEmpty == false ? trimmedValue : nil + } +} diff --git a/Sources/MusicAssistantCore/Services/SongGeneration/SongGenerationContext.swift b/Sources/MusicAssistantCore/Services/SongGeneration/SongGenerationContext.swift new file mode 100644 index 0000000..7c36858 --- /dev/null +++ b/Sources/MusicAssistantCore/Services/SongGeneration/SongGenerationContext.swift @@ -0,0 +1,171 @@ +import Foundation + +public struct SongGenerationContext: Codable, Equatable, Sendable { + public static let empty = SongGenerationContext(selectedInstruments: []) + + public var selectedInstruments: [SelectedInstrumentContext] + + public init(selectedInstruments: [SelectedInstrumentContext]) { + self.selectedInstruments = selectedInstruments + } + + public init(project: SongProject) { + let sectionsByID = Dictionary(uniqueKeysWithValues: project.sections.map { ($0.id, $0) }) + selectedInstruments = project.selectedInstrumentTracks.map { track in + SelectedInstrumentContext( + instrumentId: track.instrumentId, + displayName: Self.displayName(for: track), + playingStyle: Self.normalized(track.playingStyle), + role: Self.normalized(track.role), + autoArrangementEnabled: track.autoArrangementEnabled, + placements: track.placements.map { placement in + SelectedInstrumentPlacementContext( + sectionId: placement.sectionId, + sectionTitle: placement.sectionId.flatMap { sectionsByID[$0]?.title }, + sectionType: placement.sectionId.flatMap { sectionsByID[$0]?.type }, + startTime: placement.startTime, + endTime: placement.endTime, + direction: Self.normalized(placement.direction) + ) + } + ) + } + } + + public var hasSelectedInstruments: Bool { + !selectedInstruments.isEmpty + } + + public var sunoStyleInstrumentPhrase: String { + selectedInstruments + .map(\.stylePromptPhrase) + .joined(separator: ", ") + } + + private static func displayName(for track: InstrumentTrack) -> String { + normalized(track.variant) ?? track.instrumentId + } + + private static func normalized(_ value: String?) -> String? { + let trimmedValue = value?.trimmingCharacters(in: .whitespacesAndNewlines) + return trimmedValue?.isEmpty == false ? trimmedValue : nil + } +} + +public struct SelectedInstrumentContext: Codable, Equatable, Sendable { + public var instrumentId: String + public var displayName: String + public var playingStyle: String? + public var role: String? + public var autoArrangementEnabled: Bool + public var placements: [SelectedInstrumentPlacementContext] + + public init( + instrumentId: String, + displayName: String, + playingStyle: String? = nil, + role: String? = nil, + autoArrangementEnabled: Bool = true, + placements: [SelectedInstrumentPlacementContext] = [] + ) { + self.instrumentId = instrumentId + self.displayName = displayName + self.playingStyle = playingStyle + self.role = role + self.autoArrangementEnabled = autoArrangementEnabled + self.placements = placements + } + + public var stylePromptPhrase: String { + var details: [String] = [] + + if let role { + details.append(role) + } + + if let playingStyle { + details.append(playingStyle) + } + + let timingPhrase = placements + .map(\.stylePromptPhrase) + .filter { !$0.isEmpty } + .joined(separator: "; ") + + if !timingPhrase.isEmpty { + details.append(timingPhrase) + } + + guard !details.isEmpty else { + return displayName + } + + return "\(displayName) (\(details.joined(separator: ", ")))" + } +} + +public struct SelectedInstrumentPlacementContext: Codable, Equatable, Sendable { + public var sectionId: String? + public var sectionTitle: String? + public var sectionType: SongSectionType? + public var startTime: TimeInterval? + public var endTime: TimeInterval? + public var direction: String? + + public init( + sectionId: String? = nil, + sectionTitle: String? = nil, + sectionType: SongSectionType? = nil, + startTime: TimeInterval? = nil, + endTime: TimeInterval? = nil, + direction: String? = nil + ) { + self.sectionId = sectionId + self.sectionTitle = sectionTitle + self.sectionType = sectionType + self.startTime = startTime + self.endTime = endTime + self.direction = direction + } + + public var stylePromptPhrase: String { + var parts: [String] = [] + + if let sectionTitle, !sectionTitle.isEmpty { + parts.append(sectionTitle) + } else if let sectionType { + parts.append(sectionType.rawValue) + } + + if let timing = timingPhrase { + parts.append(timing) + } + + if let direction, !direction.isEmpty { + parts.append(direction) + } + + return parts.joined(separator: " ") + } + + private var timingPhrase: String? { + switch (startTime, endTime) { + case let (start?, end?): + return "\(Self.formattedTime(start))-\(Self.formattedTime(end))" + case let (start?, nil): + return "from \(Self.formattedTime(start))" + case let (nil, end?): + return "until \(Self.formattedTime(end))" + case (nil, nil): + return nil + } + } + + private static func formattedTime(_ time: TimeInterval) -> String { + let roundedTime = time.rounded() + if roundedTime == time { + return "\(Int(roundedTime))s" + } + return String(format: "%.1fs", time) + } +} diff --git a/Tests/MusicAssistantCoreTests/LocalSongProjectStoreTests.swift b/Tests/MusicAssistantCoreTests/LocalSongProjectStoreTests.swift index 1819403..d4c4901 100644 --- a/Tests/MusicAssistantCoreTests/LocalSongProjectStoreTests.swift +++ b/Tests/MusicAssistantCoreTests/LocalSongProjectStoreTests.swift @@ -48,6 +48,35 @@ final class LocalSongProjectStoreTests: XCTestCase { XCTAssertEqual(updatedProject, project) } + func testPersistsSelectedInstrumentsOnSongProject() async throws { + let store = LocalSongProjectStore(directoryURL: temporaryDirectoryURL) + var project = SongProject( + id: "instrument-selection-project", + title: "Instrument Selection", + idea: "Persist selected catalog instruments" + ) + + project.setInstrumentSelected(id: "oud", isSelected: true, variant: "Oud") + project.setInstrumentSelected(id: "violin", isSelected: true, variant: "Violin") + project.setInstrumentSelected(id: "oud", isSelected: false) + + try await store.create(project) + + var openedProject = try await store.open(id: project.id) + XCTAssertFalse(openedProject.isInstrumentSelected(id: "oud")) + XCTAssertTrue(openedProject.isInstrumentSelected(id: "violin")) + XCTAssertEqual(openedProject.instrumentTrack(for: "oud")?.variant, "Oud") + XCTAssertEqual(openedProject.selectedInstrumentIDs, ["violin"]) + + openedProject.setInstrumentSelected(id: "oud", isSelected: true, variant: "Oud") + try await store.save(openedProject) + + let savedProject = try await store.open(id: project.id) + XCTAssertTrue(savedProject.isInstrumentSelected(id: "oud")) + XCTAssertTrue(savedProject.isInstrumentSelected(id: "violin")) + XCTAssertEqual(savedProject.selectedInstrumentIDs, ["oud", "violin"]) + } + func testLoadsProjectListSortedByMostRecentUpdate() async throws { let store = LocalSongProjectStore(directoryURL: temporaryDirectoryURL) let older = SongProject( diff --git a/Tests/MusicAssistantCoreTests/SongGenerationContextTests.swift b/Tests/MusicAssistantCoreTests/SongGenerationContextTests.swift new file mode 100644 index 0000000..30ec710 --- /dev/null +++ b/Tests/MusicAssistantCoreTests/SongGenerationContextTests.swift @@ -0,0 +1,78 @@ +import MusicAssistantCore +import XCTest + +final class SongGenerationContextTests: XCTestCase { + func testContextIncludesOnlySelectedInstrumentsWithArrangementDetails() { + let project = SongProject( + title: "Instrument Context", + idea: "Use selected instruments in generation", + sections: [ + SongSection(id: "intro", type: .intro, title: "Opening"), + SongSection(id: "chorus", type: .chorus, title: "Final Chorus", order: 1) + ], + instruments: [ + InstrumentTrack( + instrumentId: "oud", + selected: true, + variant: "Arabic Oud", + playingStyle: "tremolo", + role: "lead motif", + autoArrangementEnabled: true, + placements: [ + InstrumentPlacement( + sectionId: "intro", + startTime: 0, + endTime: 12, + direction: "solo opening" + ) + ] + ), + InstrumentTrack( + instrumentId: "violin", + selected: false, + variant: "Violin", + role: "deselected counterline" + ) + ] + ) + + let context = SongGenerationContext(project: project) + + XCTAssertTrue(context.hasSelectedInstruments) + XCTAssertEqual(context.selectedInstruments.map(\.instrumentId), ["oud"]) + XCTAssertEqual(context.selectedInstruments.first?.displayName, "Arabic Oud") + XCTAssertEqual(context.selectedInstruments.first?.playingStyle, "tremolo") + XCTAssertEqual(context.selectedInstruments.first?.role, "lead motif") + XCTAssertEqual(context.selectedInstruments.first?.autoArrangementEnabled, true) + XCTAssertEqual(context.selectedInstruments.first?.placements.first?.sectionTitle, "Opening") + XCTAssertEqual(context.selectedInstruments.first?.placements.first?.sectionType, .intro) + XCTAssertEqual( + context.sunoStyleInstrumentPhrase, + "Arabic Oud (lead motif, tremolo, Opening 0s-12s solo opening)" + ) + } + + func testAIRequestsCarrySongGenerationContextFromProjects() { + let project = SongProject( + title: "Context Request", + idea: "Carry selected instruments", + instruments: [ + InstrumentTrack(instrumentId: "qanun", selected: true, variant: "Qanun"), + InstrumentTrack(instrumentId: "piano", selected: false, variant: "Piano") + ] + ) + + let generationRequest = SongProjectGenerationRequest( + context: AIRequestContext(userInstruction: "Generate."), + seedProject: project + ) + let updateRequest = SongProjectUpdateRequest( + context: AIRequestContext(userInstruction: "Arrange."), + project: project, + allowedScopes: [.arrangement] + ) + + XCTAssertEqual(generationRequest.songGenerationContext.selectedInstruments.map(\.instrumentId), ["qanun"]) + XCTAssertEqual(updateRequest.songGenerationContext.selectedInstruments.map(\.displayName), ["Qanun"]) + } +} diff --git a/Tests/MusicAssistantCoreTests/SongProjectPromptCompilerTests.swift b/Tests/MusicAssistantCoreTests/SongProjectPromptCompilerTests.swift new file mode 100644 index 0000000..d30fa2d --- /dev/null +++ b/Tests/MusicAssistantCoreTests/SongProjectPromptCompilerTests.swift @@ -0,0 +1,65 @@ +import MusicAssistantCore +import XCTest + +final class SongProjectPromptCompilerTests: XCTestCase { + func testCompilerUsesSelectedInstrumentsInSunoStylePrompt() throws { + let project = SongProject( + title: "Compiled Song", + idea: "Prepare Suno output", + genres: [ + GenreStyle(id: "arabic-pop", name: "Arabic Pop"), + GenreStyle(id: "cinematic", name: "Cinematic") + ], + moods: [MoodTag(id: "hopeful", name: "Hopeful")], + bpm: ManualAutoValue(mode: .manual, value: 96), + maqam: ManualAutoValue(mode: .manual, value: "Hijaz"), + sections: [ + SongSection(id: "chorus", type: .chorus, title: "Chorus", lyrics: "Section chorus") + ], + instruments: [ + InstrumentTrack( + instrumentId: "oud", + selected: true, + variant: "Arabic Oud", + playingStyle: "picked", + role: "main hook", + placements: [InstrumentPlacement(sectionId: "chorus", direction: "answer the vocal")] + ), + InstrumentTrack( + instrumentId: "drum-kit", + selected: false, + variant: "Drum Kit" + ) + ], + vocalists: [Vocalist(id: "lead", label: "Lead", voiceType: "warm tenor")], + lyrics: Lyrics(text: "Approved lyrics"), + productionDirections: [ProductionDirection(id: "lift", text: "wide chorus lift")] + ) + + let output = try SongProjectPromptCompiler().compile(project: project) + + XCTAssertEqual(output.lyricsText, "Approved lyrics") + XCTAssertTrue(output.stylePrompt.contains("Arabic Pop + Cinematic")) + XCTAssertTrue(output.stylePrompt.contains("Arabic Oud (main hook, picked, Chorus answer the vocal)")) + XCTAssertTrue(output.stylePrompt.contains("Lead warm tenor")) + XCTAssertTrue(output.stylePrompt.contains("96 BPM")) + XCTAssertTrue(output.stylePrompt.contains("maqam Hijaz")) + XCTAssertTrue(output.stylePrompt.contains("wide chorus lift")) + XCTAssertFalse(output.stylePrompt.contains("Drum Kit")) + } + + func testCompilerFallsBackToOrderedSectionLyrics() throws { + let project = SongProject( + title: "Section Lyrics", + idea: "Compile sections", + sections: [ + SongSection(id: "chorus", type: .chorus, title: "Chorus", order: 1, lyrics: "Hook line"), + SongSection(id: "verse", type: .verse, title: "Verse", order: 0, lyrics: "Verse line") + ] + ) + + let output = try SongProjectPromptCompiler().compile(project: project) + + XCTAssertEqual(output.lyricsText, "[Verse]\nVerse line\n\n[Chorus]\nHook line") + } +} diff --git a/docs/TASKS.md b/docs/TASKS.md index 59a2c43..a9a6023 100644 --- a/docs/TASKS.md +++ b/docs/TASKS.md @@ -86,8 +86,8 @@ requirement is missing and blocks implementation, record it in - [x] Add browsing/filtering by family/category. - [x] Add browsing/filtering by region/origin where useful. - [x] Add checkbox-based multi-select and deselect behavior. -- [ ] Persist selected instruments on the current Song Project. -- [ ] Make selected instruments available to OpenAI/song-generation +- [x] Persist selected instruments on the current Song Project. +- [x] Make selected instruments available to OpenAI/song-generation logic for arrangement, roles, entry/exit timing, relevant structure decisions and Suno Style Prompt generation. - [x] Preserve existing Manual/Auto arrangement behavior. @@ -96,7 +96,7 @@ requirement is missing and blocks implementation, record it in ## Phase 7 --- Prompt Compiler -- [ ] Create deterministic compiler from approved SongProject → Suno +- [x] Create deterministic compiler from approved SongProject → Suno output. - [ ] Generate lyrics text with section/performance directives where appropriate.