feat: add track-aware tool targeting and new agent tools, lifted the requirement of selection a MIDI region to use AI Agent

Introduces a shared `toolTargeting.ts` module that resolves the active MIDI
region/track from the user's current selection, enabling tools to operate on
the correct target without requiring explicit track IDs in most cases.

Adds two new read-only tools — `list_all_tracks` and
`get_user_selected_music_range_and_track` — so the agent can inspect
available tracks and the current selection context before acting.

Also refactors `AddNotesTool` to auto-create MIDI regions when no region
exists, refactors `RemoveNotesTool`, `ReadMusicTool`, and
`ReadChordProgressionTool` to use the new targeting helpers, adds a
`buildToolHistoryContent` hook to `BaseTool` for cleaner chat history
display, and updates system prompts and tests throughout.
This commit is contained in:
Xiaohan-Tian
2026-06-04 19:29:37 -07:00
parent 5c5a07b839
commit 8ee7ddac77
27 changed files with 1812 additions and 855 deletions
+149
View File
@@ -0,0 +1,149 @@
import { beforeEach, describe, expect, it, vi } from 'vitest';
import { SystemPrompts } from './SystemPrompts';
import { KGProject } from '../../core/KGProject';
import { KGMidiTrack } from '../../core/track/KGMidiTrack';
import { KGMidiRegion } from '../../core/region/KGMidiRegion';
import { KGGlobalTrack, GlobalTrackType } from '../../core/global-track/KGGlobalTrack';
import { KGMarkerRegion } from '../../core/region/KGMarkerRegion';
const storeState = {
activeRegionId: null as string | null,
selectedRegionIds: [] as string[],
};
const configState = new Map<string, unknown>();
const configManagerMock = {
getIsInitialized: vi.fn(() => true),
initialize: vi.fn().mockResolvedValue(undefined),
get: vi.fn((key: string) => configState.get(key)),
};
const coreState = {
project: new KGProject(),
selectedItems: [] as unknown[],
};
vi.mock('../../stores/projectStore', () => ({
useProjectStore: {
getState: () => storeState,
},
}));
vi.mock('../../core/config/ConfigManager', () => ({
ConfigManager: {
instance: () => configManagerMock,
},
}));
vi.mock('../../core/KGCore', () => ({
KGCore: {
instance: () => ({
getCurrentProject: () => coreState.project,
getSelectedItems: () => coreState.selectedItems,
}),
},
}));
function createProject(): KGProject {
const midiTrack = new KGMidiTrack('Piano', 1, 'acoustic_grand_piano');
const regionA = new KGMidiRegion('midi-a', 'track-1', 0, 'A', 4, 8);
const regionB = new KGMidiRegion('midi-b', 'track-1', 0, 'B', 20, 4);
midiTrack.setRegions([regionA, regionB]);
const markerTrack = new KGGlobalTrack('global-marker', 0, GlobalTrackType.Marker, 'Marker');
const markerRegion = new KGMarkerRegion('global-a', 'global-marker', 0, 'Marker A', 2, 2);
markerTrack.setRegions([markerRegion]);
const project = new KGProject('Test Project');
project.setTracks([midiTrack]);
const globalTracks = project.getGlobalTracks();
const updatedGlobalTracks = globalTracks.map(track => (
track.getType() === GlobalTrackType.Marker ? markerTrack : track
));
project.setGlobalTracks(updatedGlobalTracks);
return project;
}
describe('SystemPrompts', () => {
beforeEach(() => {
storeState.activeRegionId = null;
storeState.selectedRegionIds = [];
coreState.project = createProject();
coreState.selectedItems = [];
configState.clear();
configManagerMock.getIsInitialized.mockReturnValue(true);
configManagerMock.initialize.mockClear();
configManagerMock.get.mockClear();
SystemPrompts.clearCache();
vi.stubGlobal('fetch', vi.fn(async (input: string | URL | Request) => {
const url = String(input);
if (url.endsWith('prompts/system.md')) {
return {
ok: true,
status: 200,
text: async () => 'SYSTEM\n- BPM: {bpm}\n- Instrument: {track_instrument}',
};
}
if (url.endsWith('prompts/user_msg_appendix.md')) {
return {
ok: true,
status: 200,
text: async () => 'APPENDIX\n{selected_music_range_section}',
};
}
return {
ok: false,
status: 404,
text: async () => '',
};
}));
});
it('uses loop bounds when loop mode is enabled', async () => {
coreState.project.setIsLooping(true);
coreState.project.setLoopingRange([2, 5]);
const prompt = await SystemPrompts.getSystemPromptWithContext();
expect(prompt).toContain('APPENDIX');
expect(prompt).toContain('- Start Beat: 8');
expect(prompt).toContain('- End Beat: 24');
});
it('uses the earliest start and latest end across multiple selected regions', async () => {
storeState.selectedRegionIds = ['midi-a', 'midi-b'];
const prompt = await SystemPrompts.getSystemPromptWithContext();
expect(prompt).toContain('- Start Beat: 4');
expect(prompt).toContain('- End Beat: 24');
});
it('includes mixed regular and global selected regions in the music range span', async () => {
storeState.selectedRegionIds = ['midi-b', 'global-a'];
const prompt = await SystemPrompts.getSystemPromptWithContext();
expect(prompt).toContain('- Start Beat: 2');
expect(prompt).toContain('- End Beat: 24');
});
it('renders an explicit absence message when no music range is selected', async () => {
const prompt = await SystemPrompts.getSystemPromptWithContext();
expect(prompt).toContain('- No selected music range.');
});
it('applies appendix context to the system prompt template', async () => {
const rendered = await SystemPrompts.getPromptWithContext('Range\n{selected_music_range_section}');
expect(rendered).toContain('Range');
expect(rendered).toContain('- No selected music range.');
});
});
+46 -90
View File
@@ -1,10 +1,14 @@
import { KGCore } from '../../core/KGCore';
import { KGRegion } from '../../core/region/KGRegion';
import { KGMidiTrack } from '../../core/track/KGMidiTrack';
import { KGTrack } from '../../core/track/KGTrack';
import { useProjectStore } from '../../stores/projectStore';
import { ConfigManager } from '../../core/config/ConfigManager';
import { FLUIDR3_INSTRUMENT_MAP } from '../../constants/generalMidiConstants';
import {
findRegionById,
findRegularTrackByRegion,
resolveSelectedMusicRangeContext,
} from '../tools/toolTargeting';
/**
* Context data structure for system prompt template replacement
@@ -14,8 +18,7 @@ interface SystemPromptContext {
time_signature: string;
key_signature: string;
track_instrument: string;
current_region_start: number;
current_region_end: number;
selected_music_range_section: string;
}
/**
@@ -28,7 +31,10 @@ export class SystemPrompts {
/**
* Load the system prompt template from the public folder
*/
private static async loadTemplate(templatePath: string = 'prompts/system.md'): Promise<string> {
private static async loadTemplate(
templatePath: string = 'prompts/system.md',
fallbackContent: string = this.FALLBACK_PROMPT,
): Promise<string> {
if (this.cachedTemplates.has(templatePath)) {
return this.cachedTemplates.get(templatePath)!;
}
@@ -44,40 +50,10 @@ export class SystemPrompts {
return template;
} catch (error) {
console.error('Failed to load system prompt template:', error);
return this.FALLBACK_PROMPT;
return fallbackContent;
}
}
/**
* Find a region by ID across all tracks
*/
private static findRegionById(regionId: string): KGRegion | null {
const core = KGCore.instance();
const project = core.getCurrentProject();
const tracks = project.getTracks();
for (const track of tracks) {
const regions = track.getRegions();
const region = regions.find(r => r.getId() === regionId);
if (region) {
return region;
}
}
return null;
}
/**
* Find track that contains the given region
*/
private static findTrackByRegion(region: KGRegion): KGTrack | null {
const core = KGCore.instance();
const project = core.getCurrentProject();
const tracks = project.getTracks();
return tracks.find(track => track.getRegions().includes(region)) || null;
}
/**
* Extract current project context from KGCore
*/
@@ -98,9 +74,9 @@ export class SystemPrompts {
// Step 1: Check if there's an active piano roll region
const activeRegionId = this.getActiveRegionId();
if (activeRegionId) {
const activeRegion = this.findRegionById(activeRegionId);
const activeRegion = findRegionById(activeRegionId);
if (activeRegion) {
const track = this.findTrackByRegion(activeRegion);
const track = findRegularTrackByRegion(activeRegion);
if (track && track instanceof KGMidiTrack) {
trackInstrument = FLUIDR3_INSTRUMENT_MAP[track.getInstrument()].displayName;
}
@@ -108,10 +84,12 @@ export class SystemPrompts {
} else {
// Step 2: Check if user has selected region(s)
const selectedItems = core.getSelectedItems();
const selectedRegion = selectedItems.find(item => item instanceof KGRegion) as KGRegion;
const selectedRegion = selectedItems.find((item): item is KGRegion => (
item instanceof KGRegion && findRegularTrackByRegion(item) !== null
)) ?? null;
if (selectedRegion) {
const track = this.findTrackByRegion(selectedRegion);
const track = findRegularTrackByRegion(selectedRegion);
if (track && track instanceof KGMidiTrack) {
trackInstrument = FLUIDR3_INSTRUMENT_MAP[track.getInstrument()].displayName;
}
@@ -132,54 +110,26 @@ export class SystemPrompts {
/**
* Get active region ID from project store
*/
private static getActiveRegionId(): string | null {
private static getStoreState(): ReturnType<typeof useProjectStore.getState> | null {
try {
const store = useProjectStore.getState();
return store.activeRegionId;
return useProjectStore.getState();
} catch {
return null;
}
}
/**
* Extract current region context with fallback logic
* Get active region ID from project store
*/
private static extractRegionContext(): Partial<SystemPromptContext> {
const core = KGCore.instance();
const project = core.getCurrentProject();
// Step 1: Try active piano roll region
const activeRegionId = this.getActiveRegionId();
if (activeRegionId) {
const activeRegion = this.findRegionById(activeRegionId);
if (activeRegion) {
return {
current_region_start: activeRegion.getStartFromBeat(),
current_region_end: activeRegion.getStartFromBeat() + activeRegion.getLength(),
};
}
}
// Step 2: Try selected region
const selectedItems = core.getSelectedItems();
const selectedRegion = selectedItems.find(item => item instanceof KGRegion) as KGRegion;
if (selectedRegion) {
return {
current_region_start: selectedRegion.getStartFromBeat(),
current_region_end: selectedRegion.getStartFromBeat() + selectedRegion.getLength(),
};
}
// Step 3: Fallback to project bounds
const timeSignature = project.getTimeSignature();
const beatsPerBar = timeSignature.numerator;
const maxBars = project.getMaxBars();
return {
current_region_start: 0,
current_region_end: maxBars * beatsPerBar,
};
private static getActiveRegionId(): string | null {
return this.getStoreState()?.activeRegionId ?? null;
}
/**
* Build the selected music range section for prompt templates.
*/
private static buildSelectedMusicRangeSection(): string {
return resolveSelectedMusicRangeContext().section;
}
/**
@@ -187,15 +137,13 @@ export class SystemPrompts {
*/
private static getFullContext(): SystemPromptContext {
const projectContext = this.extractProjectContext();
const regionContext = this.extractRegionContext();
return {
bpm: projectContext.bpm || 120,
time_signature: projectContext.time_signature || '4/4',
key_signature: projectContext.key_signature || 'C major',
track_instrument: projectContext.track_instrument || 'Piano',
current_region_start: regionContext.current_region_start || 0,
current_region_end: regionContext.current_region_end || 32,
bpm: projectContext.bpm ?? 120,
time_signature: projectContext.time_signature ?? '4/4',
key_signature: projectContext.key_signature ?? 'C major',
track_instrument: projectContext.track_instrument ?? 'Piano',
selected_music_range_section: this.buildSelectedMusicRangeSection(),
};
}
@@ -210,8 +158,7 @@ export class SystemPrompts {
result = result.replace(/{time_signature}/g, context.time_signature);
result = result.replace(/{key_signature}/g, context.key_signature);
result = result.replace(/{track_instrument}/g, context.track_instrument);
result = result.replace(/{current_region_start}/g, context.current_region_start.toString());
result = result.replace(/{current_region_end}/g, context.current_region_end.toString());
result = result.replace(/{selected_music_range_section}/g, context.selected_music_range_section);
return result;
}
@@ -234,8 +181,17 @@ export class SystemPrompts {
*/
static async getSystemPromptWithContext(templatePath?: string): Promise<string> {
try {
const template = await this.loadTemplate(templatePath);
let promptWithContext = await this.getPromptWithContext(template);
const context = this.getFullContext();
const [template, appendixTemplate] = await Promise.all([
this.loadTemplate(templatePath),
this.loadTemplate('prompts/user_msg_appendix.md', ''),
]);
let promptWithContext = this.replaceTemplateVariables(template, context);
const appendixWithContext = this.replaceTemplateVariables(appendixTemplate, context);
if (appendixWithContext.trim().length > 0) {
promptWithContext += `\n\n${appendixWithContext}`;
}
// Append custom instructions from config if provided
try {
+141 -53
View File
@@ -1,9 +1,15 @@
import { beforeEach, describe, expect, it, vi } from 'vitest';
import { AddNotesTool } from './AddNotesTool';
import {
NO_MIDI_TARGET_HISTORY_MESSAGE,
NO_MIDI_TARGET_RAW_MESSAGE,
NO_MIDI_TARGET_UI_MESSAGE,
} from './toolTargeting';
import { KGCore } from '../../core/KGCore';
import { KGProject } from '../../core/KGProject';
import { KGMidiTrack } from '../../core/track/KGMidiTrack';
import { KGMidiRegion } from '../../core/region/KGMidiRegion';
import { KGMidiNote } from '../../core/midi/KGMidiNote';
const storeState = {
activeRegionId: null as string | null,
@@ -15,82 +21,164 @@ vi.mock('../../stores/projectStore', () => ({
},
}));
function mockCore(project: KGProject, selectedItems: unknown[] = []) {
vi.spyOn(KGCore, 'instance').mockReturnValue({
getCurrentProject: () => project,
getSelectedItems: () => selectedItems,
executeCommand: (command: { execute(): void }) => command.execute(),
} as unknown as KGCore);
}
describe('AddNotesTool', () => {
beforeEach(() => {
storeState.activeRegionId = null;
vi.restoreAllMocks();
});
it('builds a compact summary for successful note creation', () => {
it('builds summaries for an active MIDI region target', () => {
const track = new KGMidiTrack('Lead', 1);
const region = new KGMidiRegion('region-1', track.getId().toString(), track.getTrackIndex(), 'Verse Melody', 0, 32);
track.setRegions([region]);
const project = new KGProject('summary-project', 8, 0, 120, { numerator: 4, denominator: 4 }, 'C major');
project.setTracks([track]);
storeState.activeRegionId = region.getId();
vi.spyOn(KGCore, 'instance').mockReturnValue({
getCurrentProject: () => project,
getSelectedItems: () => [],
} as unknown as KGCore);
mockCore(project);
const tool = new AddNotesTool();
const summary = tool.buildToolResultDisplayContent(
{
notes: [
{ pitch: 'C4', start: 16, length: 4 },
{ pitch: 'E4', start: 20, length: 8 },
],
},
{ success: true, result: 'raw result' },
);
expect(summary).toBe(
'Successfully created 2 notes in region **Verse Melody** on track **Lead**, spanning bars 5 to 7.'
);
});
it('builds a confirmation summary for note creation', () => {
const track = new KGMidiTrack('Lead', 1);
const region = new KGMidiRegion('region-1', track.getId().toString(), track.getTrackIndex(), 'Verse Melody', 0, 32);
track.setRegions([region]);
const project = new KGProject('summary-project', 8, 0, 120, { numerator: 4, denominator: 4 }, 'C major');
project.setTracks([track]);
storeState.activeRegionId = region.getId();
vi.spyOn(KGCore, 'instance').mockReturnValue({
getCurrentProject: () => project,
getSelectedItems: () => [],
} as unknown as KGCore);
const tool = new AddNotesTool();
expect(tool.isReadOnlyTool()).toBe(false);
expect(tool.buildConfirmationContent({
const args = {
notes: [
{ pitch: 'C4', start: 16, length: 4 },
{ pitch: 'E4', start: 20, length: 8 },
],
})).toBe(
'Allow creating 2 notes in region **Verse Melody** on track **Lead**, spanning bars 5 to 7?'
};
expect(tool.buildToolResultDisplayContent(args, { success: true, result: 'raw result' })).toBe(
'Successfully created 2 notes in region **Verse Melody** on track **Lead**, spanning bars 5 to 7.',
);
expect(tool.buildConfirmationContent(args)).toBe(
'Allow creating 2 notes on track **Lead** in region **Verse Melody**, spanning bars 5 to 7?',
);
});
it('returns no compact summary when the target region cannot be resolved', () => {
const project = new KGProject('summary-project', 8, 0, 120, { numerator: 4, denominator: 4 }, 'C major');
vi.spyOn(KGCore, 'instance').mockReturnValue({
getCurrentProject: () => project,
getSelectedItems: () => [],
} as unknown as KGCore);
it('creates a new MIDI region on the requested track when no region overlaps', async () => {
const track = new KGMidiTrack('Lead', 1);
const existingRegion = new KGMidiRegion('region-1', track.getId().toString(), track.getTrackIndex(), 'Intro', 0, 4);
track.setRegions([existingRegion]);
const project = new KGProject('create-region-project', 8, 0, 120, { numerator: 4, denominator: 4 }, 'C major');
project.setTracks([track]);
mockCore(project);
const tool = new AddNotesTool();
const summary = tool.buildToolResultDisplayContent(
{
notes: [{ pitch: 'C4', start: 0, length: 4 }],
},
{ success: true, result: 'raw result' },
);
const result = await tool.execute({
track_id: track.getId().toString(),
notes: [
{ pitch: 'C4', start: 16, length: 2 },
{ pitch: 'E4', start: 18, length: 2 },
],
});
expect(summary).toBeUndefined();
expect(result.success).toBe(true);
expect(track.getRegions()).toHaveLength(2);
const createdRegion = track.getRegions()[1] as KGMidiRegion;
expect(createdRegion.getName()).toBe('Lead Region');
expect(createdRegion.getStartFromBeat()).toBe(16);
expect(createdRegion.getLength()).toBe(4);
expect(createdRegion.getNotes()).toHaveLength(2);
expect(createdRegion.getNotes().map(note => note.getStartBeat())).toEqual([0, 2]);
});
it('targets a track by track_name when track_id is omitted', async () => {
const targetTrack = new KGMidiTrack('Lead', 1);
const otherTrack = new KGMidiTrack('Bass', 2);
targetTrack.setRegions([new KGMidiRegion('region-1', targetTrack.getId().toString(), targetTrack.getTrackIndex(), 'Intro', 0, 4)]);
const project = new KGProject('track-name-project', 8, 0, 120, { numerator: 4, denominator: 4 }, 'C major');
project.setTracks([targetTrack, otherTrack]);
mockCore(project);
const tool = new AddNotesTool();
const result = await tool.execute({
track_name: 'Lead',
notes: [{ pitch: 'C4', start: 16, length: 2 }],
});
expect(result.success).toBe(true);
expect(targetTrack.getRegions()).toHaveLength(2);
expect(otherTrack.getRegions()).toHaveLength(0);
});
it('uses track_id when both track_id and track_name are provided', async () => {
const leadTrack = new KGMidiTrack('Lead', 1);
const bassTrack = new KGMidiTrack('Bass', 2);
const project = new KGProject('track-id-precedence-project', 8, 0, 120, { numerator: 4, denominator: 4 }, 'C major');
project.setTracks([leadTrack, bassTrack]);
mockCore(project);
const tool = new AddNotesTool();
const result = await tool.execute({
track_id: bassTrack.getId().toString(),
track_name: 'Lead',
notes: [{ pitch: 'C4', start: 8, length: 2 }],
});
expect(result.success).toBe(true);
expect(leadTrack.getRegions()).toHaveLength(0);
expect(bassTrack.getRegions()).toHaveLength(1);
expect((bassTrack.getRegions()[0] as KGMidiRegion).getName()).toBe('Bass Region');
});
it('uses the first matching track when duplicate track names exist', async () => {
const firstLead = new KGMidiTrack('Lead', 1);
const secondLead = new KGMidiTrack('Lead', 2);
const project = new KGProject('duplicate-track-name-project', 8, 0, 120, { numerator: 4, denominator: 4 }, 'C major');
project.setTracks([firstLead, secondLead]);
mockCore(project);
const tool = new AddNotesTool();
const result = await tool.execute({
track_name: 'Lead',
notes: [{ pitch: 'C4', start: 4, length: 2 }],
});
expect(result.success).toBe(true);
expect(firstLead.getRegions()).toHaveLength(1);
expect(secondLead.getRegions()).toHaveLength(0);
});
it('chooses the largest overlapping region and auto-expands it to fit the notes', async () => {
const track = new KGMidiTrack('Lead', 1);
const regionA = new KGMidiRegion('region-a', track.getId().toString(), track.getTrackIndex(), 'Region A', 0, 4);
const regionB = new KGMidiRegion('region-b', track.getId().toString(), track.getTrackIndex(), 'Region B', 4, 4);
regionB.setNotes([new KGMidiNote('note-existing', 0, 1, 60, 100)]);
track.setRegions([regionA, regionB]);
const project = new KGProject('expand-region-project', 8, 0, 120, { numerator: 4, denominator: 4 }, 'C major');
project.setTracks([track]);
mockCore(project);
const tool = new AddNotesTool();
const result = await tool.execute({
track_id: track.getId().toString(),
notes: [{ pitch: 'G4', start: 2, length: 5 }],
});
expect(result.success).toBe(true);
expect(regionA.getNotes()).toHaveLength(0);
expect(regionB.getStartFromBeat()).toBe(2);
expect(regionB.getLength()).toBe(6);
expect(regionB.getNotes()).toHaveLength(2);
expect(regionB.getNotes().find(note => note.getId() === 'note-existing')?.getStartBeat()).toBe(2);
expect(regionB.getNotes().find(note => note.getId() !== 'note-existing')?.getStartBeat()).toBe(0);
});
it('returns distinct raw, history, and UI guidance when no MIDI target is available', async () => {
const project = new KGProject('no-target-project', 8, 0, 120, { numerator: 4, denominator: 4 }, 'C major');
mockCore(project);
const tool = new AddNotesTool();
const args = { notes: [{ pitch: 'C4', start: 0, length: 1 }] };
const result = await tool.execute(args);
expect(result).toEqual({ success: false, result: NO_MIDI_TARGET_RAW_MESSAGE });
expect(tool.buildToolHistoryContent(args, result)).toBe(NO_MIDI_TARGET_HISTORY_MESSAGE);
expect(tool.buildToolResultDisplayContent(args, result)).toBe(NO_MIDI_TARGET_UI_MESSAGE);
});
});
+312 -168
View File
@@ -1,10 +1,28 @@
import { BaseTool } from './BaseTool';
import type { ToolResult, ToolParameter } from './BaseTool';
import {
NO_MIDI_TARGET_HISTORY_MESSAGE,
NO_MIDI_TARGET_RAW_MESSAGE,
NO_MIDI_TARGET_UI_MESSAGE,
getTrackDisplayName,
resolveMidiTrackByIdOrName,
resolveActiveOrSelectedMidiRegionContext,
} from './toolTargeting';
import { CreateNotesCommand } from '../../core/commands/note/CreateNotesCommand';
import type { NoteCreationData } from '../../core/commands/note/CreateNotesCommand';
import { KGMidiRegion } from '../../core/region/KGMidiRegion';
import { useProjectStore } from '../../stores/projectStore';
import { KGCommand } from '../../core/commands/KGCommand';
import { CreateRegionCommand } from '../../core/commands/region/CreateRegionCommand';
import { ResizeRegionCommand } from '../../core/commands/region/ResizeRegionCommand';
import { KGCore } from '../../core/KGCore';
import { KGMidiRegion } from '../../core/region/KGMidiRegion';
import { KGMidiTrack } from '../../core/track/KGMidiTrack';
interface RequestedNote {
pitch: string;
start: number;
length: number;
velocity?: number;
}
interface AddNotesSummaryData {
noteCount: number;
@@ -12,15 +30,109 @@ interface AddNotesSummaryData {
trackName: string;
earliestNoteStartBar: number;
latestNoteEndBar: number;
createdRegion: boolean;
}
interface NoteSpan {
startBeat: number;
endBeat: number;
}
interface ResolvedRegionContext {
track: KGMidiTrack;
trackName: string;
regionName: string;
regionId?: string;
finalRegionStartBeat: number;
finalRegionLength: number;
createdRegion: boolean;
}
class AddNotesToResolvedRegionCommand extends KGCommand {
private readonly resolvedRegion: ResolvedRegionContext;
private readonly notes: Array<RequestedNote & { midiPitch: number; velocity: number }>;
private createRegionCommand: CreateRegionCommand | null = null;
private resizeRegionCommand: ResizeRegionCommand | null = null;
private createNotesCommand: CreateNotesCommand | null = null;
constructor(
resolvedRegion: ResolvedRegionContext,
notes: Array<RequestedNote & { midiPitch: number; velocity: number }>,
) {
super();
this.resolvedRegion = resolvedRegion;
this.notes = notes;
}
execute(): void {
let regionId = this.resolvedRegion.regionId;
if (this.resolvedRegion.createdRegion) {
this.createRegionCommand = new CreateRegionCommand(
this.resolvedRegion.track.getId().toString(),
this.resolvedRegion.track.getTrackIndex(),
this.resolvedRegion.finalRegionStartBeat,
this.resolvedRegion.finalRegionLength,
this.resolvedRegion.regionName,
);
this.createRegionCommand.execute();
regionId = this.createRegionCommand.getRegionId();
} else if (regionId) {
const existingRegion = this.findMidiRegion(regionId);
if (
existingRegion.getStartFromBeat() !== this.resolvedRegion.finalRegionStartBeat
|| existingRegion.getLength() !== this.resolvedRegion.finalRegionLength
) {
this.resizeRegionCommand = new ResizeRegionCommand(
regionId,
this.resolvedRegion.finalRegionStartBeat,
this.resolvedRegion.finalRegionLength,
);
this.resizeRegionCommand.execute();
}
}
if (!regionId) {
throw new Error('Unable to resolve the MIDI region for note creation.');
}
const noteCreationData: NoteCreationData[] = this.notes.map(note => ({
regionId,
startBeat: note.start - this.resolvedRegion.finalRegionStartBeat,
endBeat: note.start - this.resolvedRegion.finalRegionStartBeat + note.length,
pitch: note.midiPitch,
velocity: note.velocity,
}));
this.createNotesCommand = new CreateNotesCommand(noteCreationData);
this.createNotesCommand.execute();
}
undo(): void {
this.createNotesCommand?.undo();
this.resizeRegionCommand?.undo();
this.createRegionCommand?.undo();
}
getDescription(): string {
return `Add ${this.notes.length} note${this.notes.length === 1 ? '' : 's'}`;
}
private findMidiRegion(regionId: string): KGMidiRegion {
const tracks = KGCore.instance().getCurrentProject().getTracks();
for (const track of tracks) {
const region = track.getRegions().find(candidate => candidate.getId() === regionId);
if (region instanceof KGMidiRegion) {
return region;
}
}
throw new Error(`MIDI region with ID "${regionId}" not found.`);
}
}
/**
* Tool for adding notes to MIDI regions
* Integrates with the existing command system for undo/redo support
*/
export class AddNotesTool extends BaseTool {
readonly name = 'add_notes';
readonly description = 'Add one or more MIDI notes to the current region. Use this to create melodies, chords, or any musical content. Notes use absolute beat positions on the project timeline — not relative to the region start.';
readonly description = 'Add one or more MIDI notes to a target track or the currently active MIDI region. Use track_id when the user identifies a track. Regions are resolved or created automatically, so you should think in terms of tracks rather than clips. Notes use absolute beat positions on the project timeline.';
override isReadOnlyTool(): boolean {
return false;
@@ -38,44 +150,64 @@ export class AddNotesTool extends BaseTool {
pitch: {
type: 'string',
description: 'Pitch in scientific notation: note name, optional accidental (# or b), and octave number. Examples: "C4" (middle C), "F#3" (F-sharp 3rd octave), "Bb2" (B-flat 2nd octave).',
required: true
required: true,
},
start: {
type: 'number',
description: 'Start beat — the absolute beat position on the project timeline where the note begins. This is NOT relative to the region — beat 6 means beat 6 in the project regardless of where the region begins. Fractional values are supported (e.g., 0.5 = half a beat after beat 0).',
required: true
description: 'Start beat — the absolute beat position on the project timeline where the note begins. This is NOT relative to the region or clip — beat 6 means beat 6 in the project regardless of where any MIDI region begins. Fractional values are supported (e.g., 0.5 = half a beat after beat 0).',
required: true,
},
length: {
type: 'number',
description: 'Duration of the note in beats. In 4/4 time: 4 = whole note, 2 = half note, 1 = quarter note, 0.5 = eighth note, 0.25 = sixteenth note.',
required: true
required: true,
},
velocity: {
type: 'number',
description: 'Note velocity / loudness from 1 (softest) to 127 (loudest). Defaults to 127 if omitted.',
required: false
}
}
}
required: false,
},
},
},
},
region_id: {
track_id: {
type: 'string',
description: 'Target region ID. If omitted, uses the currently active piano roll region or selected region.',
required: false
}
description: 'Optional target MIDI track ID. If provided, the app automatically resolves the best overlapping MIDI region on that track for the requested note span, expands it if needed, or creates a new MIDI region when no overlap exists.',
required: false,
},
track_name: {
type: 'string',
description: 'Optional target MIDI track name. Used only when track_id is omitted. If multiple MIDI tracks share the same name, the first matching track is used.',
required: false,
},
};
buildToolResultDisplayContent(args: Record<string, unknown> | null, toolResult: ToolResult): string | undefined {
if (!toolResult.success || !args) {
override buildToolResultDisplayContent(args: Record<string, unknown> | null, toolResult: ToolResult): string | undefined {
if (!args) {
return undefined;
}
if (!toolResult.success) {
return toolResult.result === NO_MIDI_TARGET_RAW_MESSAGE ? NO_MIDI_TARGET_UI_MESSAGE : undefined;
}
const summary = this.buildSummaryData(args);
if (!summary) {
return undefined;
}
return `Successfully created ${summary.noteCount} ${summary.noteCount === 1 ? 'note' : 'notes'} in region **${summary.regionName}** on track **${summary.trackName}**, spanning bars ${summary.earliestNoteStartBar} to ${summary.latestNoteEndBar}.`;
const regionLabel = summary.createdRegion
? `new region **${summary.regionName}**`
: `region **${summary.regionName}**`;
return `Successfully created ${summary.noteCount} ${summary.noteCount === 1 ? 'note' : 'notes'} in ${regionLabel} on track **${summary.trackName}**, spanning bars ${summary.earliestNoteStartBar} to ${summary.latestNoteEndBar}.`;
}
override buildToolHistoryContent(args: Record<string, unknown> | null, toolResult: ToolResult): string | undefined {
if (!args || toolResult.success) {
return undefined;
}
return toolResult.result === NO_MIDI_TARGET_RAW_MESSAGE ? NO_MIDI_TARGET_HISTORY_MESSAGE : undefined;
}
override buildConfirmationContent(args: Record<string, unknown> | null): string | undefined {
@@ -88,222 +220,234 @@ export class AddNotesTool extends BaseTool {
return undefined;
}
return `Allow creating ${summary.noteCount} ${summary.noteCount === 1 ? 'note' : 'notes'} in region **${summary.regionName}** on track **${summary.trackName}**, spanning bars ${summary.earliestNoteStartBar} to ${summary.latestNoteEndBar}?`;
const regionVerb = summary.createdRegion ? 'a new region' : `region **${summary.regionName}**`;
return `Allow creating ${summary.noteCount} ${summary.noteCount === 1 ? 'note' : 'notes'} on track **${summary.trackName}** in ${regionVerb}, spanning bars ${summary.earliestNoteStartBar} to ${summary.latestNoteEndBar}?`;
}
async execute(params: Record<string, unknown>): Promise<ToolResult> {
try {
// Validate parameters
this.validateParameters(params);
const notes = params.notes as Array<{
pitch: string;
start: number;
length: number;
velocity?: number;
}>;
const regionId = params.region_id as string | undefined;
// Find the target region
const targetRegion = this.findTargetRegion(regionId);
if (!targetRegion) {
return this.createErrorResult(
regionId
? `Region with ID "${regionId}" not found or is not a MIDI region`
: 'No active or selected MIDI region found. Please open the piano roll with a region or select a MIDI region first.'
);
const notes = params.notes as RequestedNote[];
if (notes.length === 0) {
return this.createErrorResult('No notes were provided.');
}
// Validate and convert notes to creation data
const noteCreationData: NoteCreationData[] = [];
const createdNotes: Array<{ pitch: string; start: number; length: number }> = [];
const validatedNotes: Array<RequestedNote & { midiPitch: number; velocity: number }> = [];
for (const note of notes) {
try {
const midiPitch = this.convertPitchToMidi(note.pitch);
const velocity = note.velocity ?? 127;
// Validate velocity range
if (velocity < 1 || velocity > 127) {
return this.createErrorResult(`Invalid velocity ${velocity}. Must be between 1 and 127.`);
}
// Validate beat positions
if (note.start < 0) {
return this.createErrorResult(`Invalid start ${note.start}. Must be >= 0.`);
}
if (note.length <= 0) {
return this.createErrorResult(`Invalid length ${note.length}. Must be > 0.`);
}
// Adjust note position relative to region's start beat
const regionStartBeat = targetRegion.getStartFromBeat();
const adjustedStartBeat = note.start - regionStartBeat;
const adjustedEndBeat = adjustedStartBeat + note.length;
// Create note creation data
noteCreationData.push({
regionId: targetRegion.getId(),
startBeat: adjustedStartBeat,
endBeat: adjustedEndBeat,
pitch: midiPitch,
velocity
});
createdNotes.push({
pitch: note.pitch,
start: note.start,
length: note.length
});
validatedNotes.push({ ...note, midiPitch, velocity });
} catch (error) {
return this.createErrorResult(`Invalid note pitch "${note.pitch}": ${error}`);
}
}
// Execute the bulk note creation command
const command = new CreateNotesCommand(noteCreationData);
const trackId = params.track_id as string | undefined;
const trackName = params.track_name as string | undefined;
if (trackId || trackName) {
const explicitTrack = resolveMidiTrackByIdOrName(trackId, trackName);
if (!explicitTrack) {
return this.createErrorResult(
trackId
? `Track with ID "${trackId}" not found or is not a MIDI track.`
: `Track with name "${trackName}" not found or is not a MIDI track.`,
);
}
}
const resolvedRegion = this.resolveTargetRegion(trackId, trackName, this.getNoteSpan(validatedNotes));
if (!resolvedRegion) {
return this.createErrorResult(NO_MIDI_TARGET_RAW_MESSAGE);
}
const command = new AddNotesToResolvedRegionCommand(resolvedRegion, validatedNotes);
await this.executeCommand(command);
// Create success message
const noteCount = createdNotes.length;
const noteList = createdNotes
const noteList = validatedNotes
.map(note => `${note.pitch} (beat ${note.start}, length ${note.length})`)
.join(', ');
return this.createSuccessResult(
`Successfully created ${noteCount} note${noteCount > 1 ? 's' : ''}: ${noteList}`
);
const actionPrefix = resolvedRegion.createdRegion
? `Successfully created ${validatedNotes.length} note${validatedNotes.length > 1 ? 's' : ''} on track "${resolvedRegion.trackName}" by creating MIDI region "${resolvedRegion.regionName}"`
: `Successfully created ${validatedNotes.length} note${validatedNotes.length > 1 ? 's' : ''} in MIDI region "${resolvedRegion.regionName}" on track "${resolvedRegion.trackName}"`;
return this.createSuccessResult(`${actionPrefix}: ${noteList}`);
} catch (error) {
return this.createErrorResult(`Failed to create notes: ${error}`);
}
}
/**
* Find the target region for note creation
* Priority: 1) Specified regionId, 2) Active piano roll region, 3) Selected regions, 4) Error if none found
*/
private findTargetRegion(regionId?: string): KGMidiRegion | null {
return this.findTargetRegionContext(regionId)?.region ?? null;
}
private findTargetRegionContext(regionId?: string): { region: KGMidiRegion; trackName: string } | null {
const project = this.getCurrentProject();
const tracks = project.getTracks();
if (regionId) {
for (const track of tracks) {
const regions = track.getRegions();
const region = regions.find(r => r.getId() === regionId);
if (region && region instanceof KGMidiRegion) {
return {
region,
trackName: track.getName() || `Track ${track.getTrackIndex() + 1}`,
};
}
}
return null;
} else {
const storeState = useProjectStore.getState();
if (storeState.activeRegionId) {
for (const track of tracks) {
const regions = track.getRegions();
const region = regions.find(r => r.getId() === storeState.activeRegionId);
if (region && region instanceof KGMidiRegion) {
return {
region,
trackName: track.getName() || `Track ${track.getTrackIndex() + 1}`,
};
}
}
}
const core = this.getKGCore();
const selectedItems = core.getSelectedItems();
for (const item of selectedItems) {
if (item instanceof KGMidiRegion) {
const track = tracks.find(candidate => candidate.getId().toString() === item.getTrackId());
return {
region: item,
trackName: track?.getName() || `Track ${item.getTrackIndex() + 1}`,
};
}
}
return null;
}
}
private buildSummaryData(args: Record<string, unknown>): AddNotesSummaryData | null {
const typedArgs = args as {
notes?: Array<{ start: number; length: number }>;
region_id?: string;
track_id?: string;
track_name?: string;
};
if (!Array.isArray(typedArgs.notes) || typedArgs.notes.length === 0) {
return null;
}
const targetRegion = this.findTargetRegionContext(typedArgs.region_id);
if (!targetRegion) {
const span = this.getNoteSpan(typedArgs.notes);
const resolvedRegion = this.resolveTargetRegion(typedArgs.track_id, typedArgs.track_name, span);
if (!resolvedRegion) {
return null;
}
const beatsPerBar = this.getCurrentProject().getTimeSignature().numerator;
const earliestNoteStartBeat = Math.min(...typedArgs.notes.map(note => note.start));
const latestNoteEndBeat = Math.max(...typedArgs.notes.map(note => note.start + note.length));
return {
noteCount: typedArgs.notes.length,
regionName: targetRegion.region.getName(),
trackName: targetRegion.trackName,
earliestNoteStartBar: Math.floor(earliestNoteStartBeat / beatsPerBar) + 1,
latestNoteEndBar: Math.max(1, Math.ceil(latestNoteEndBeat / beatsPerBar)),
regionName: resolvedRegion.regionName,
trackName: resolvedRegion.trackName,
earliestNoteStartBar: Math.floor(span.startBeat / beatsPerBar) + 1,
latestNoteEndBar: Math.max(1, Math.ceil(span.endBeat / beatsPerBar)),
createdRegion: resolvedRegion.createdRegion,
};
}
/**
* Get KGCore instance for selection access
*/
private getKGCore() {
return KGCore.instance();
private resolveTargetRegion(
trackId: string | undefined,
trackName: string | undefined,
span: NoteSpan,
): ResolvedRegionContext | null {
if (trackId || trackName) {
return this.resolveTrackTarget(trackId, trackName, span);
}
const activeRegion = resolveActiveOrSelectedMidiRegionContext();
if (!activeRegion) {
return null;
}
const region = activeRegion.region;
const regionStartBeat = region.getStartFromBeat();
const regionEndBeat = regionStartBeat + region.getLength();
return {
track: activeRegion.track,
trackName: activeRegion.trackName,
regionId: region.getId(),
regionName: region.getName(),
finalRegionStartBeat: Math.min(regionStartBeat, span.startBeat),
finalRegionLength: Math.max(regionEndBeat, span.endBeat) - Math.min(regionStartBeat, span.startBeat),
createdRegion: false,
};
}
private resolveTrackTarget(
trackId: string | undefined,
trackName: string | undefined,
span: NoteSpan,
): ResolvedRegionContext | null {
const track = resolveMidiTrackByIdOrName(trackId, trackName);
if (!track) {
return null;
}
const resolvedTrackName = getTrackDisplayName(track);
const midiRegions = track.getRegions().filter(region => region instanceof KGMidiRegion) as KGMidiRegion[];
const selectedRegion = this.pickBestOverlappingRegion(midiRegions, span);
if (!selectedRegion) {
return {
track,
trackName: resolvedTrackName,
regionName: `${resolvedTrackName} Region`,
finalRegionStartBeat: span.startBeat,
finalRegionLength: span.endBeat - span.startBeat,
createdRegion: true,
};
}
const regionStartBeat = selectedRegion.getStartFromBeat();
const regionEndBeat = regionStartBeat + selectedRegion.getLength();
const finalRegionStartBeat = Math.min(regionStartBeat, span.startBeat);
const finalRegionEndBeat = Math.max(regionEndBeat, span.endBeat);
return {
track,
trackName: resolvedTrackName,
regionId: selectedRegion.getId(),
regionName: selectedRegion.getName(),
finalRegionStartBeat,
finalRegionLength: finalRegionEndBeat - finalRegionStartBeat,
createdRegion: false,
};
}
private pickBestOverlappingRegion(regions: KGMidiRegion[], span: NoteSpan): KGMidiRegion | null {
let bestRegion: KGMidiRegion | null = null;
let bestOverlap = -1;
let bestDistance = Number.POSITIVE_INFINITY;
for (const region of regions) {
const regionStart = region.getStartFromBeat();
const regionEnd = regionStart + region.getLength();
const overlap = Math.min(regionEnd, span.endBeat) - Math.max(regionStart, span.startBeat);
if (overlap <= 0) {
continue;
}
const distance = Math.abs(regionStart - span.startBeat);
if (overlap > bestOverlap || (overlap === bestOverlap && distance < bestDistance)) {
bestRegion = region;
bestOverlap = overlap;
bestDistance = distance;
}
}
return bestRegion;
}
private getNoteSpan(notes: Array<{ start: number; length: number }>): NoteSpan {
return {
startBeat: Math.min(...notes.map(note => note.start)),
endBeat: Math.max(...notes.map(note => note.start + note.length)),
};
}
/**
* Convert pitch string to MIDI note number
* Supports formats like: C4, F#3, Bb2, C#5
*/
private convertPitchToMidi(pitch: string): number {
const match = pitch.match(/^([A-G])([#b]?)(\d+)$/);
if (!match) {
throw new Error(`Invalid pitch format "${pitch}". Use format like "C4", "F#3", "Bb2"`);
}
const [, noteName, accidental, octaveStr] = match;
const octave = parseInt(octaveStr);
// Base MIDI notes for C octave (C4 = 60)
const octave = parseInt(octaveStr, 10);
const noteOffsets: Record<string, number> = {
'C': 0, 'D': 2, 'E': 4, 'F': 5, 'G': 7, 'A': 9, 'B': 11
C: 0,
D: 2,
E: 4,
F: 5,
G: 7,
A: 9,
B: 11,
};
let midiNote = (octave + 1) * 12 + noteOffsets[noteName];
// Apply accidentals
if (accidental === '#') {
midiNote += 1;
} else if (accidental === 'b') {
midiNote -= 1;
}
// Validate MIDI range
if (midiNote < 0 || midiNote > 127) {
throw new Error(`Note "${pitch}" is out of MIDI range (0-127)`);
}
return midiNote;
}
}
+11
View File
@@ -99,6 +99,17 @@ export abstract class BaseTool {
return undefined;
}
/**
* Optionally build the user-visible tool result text stored in chat history.
* Raw tool results remain the canonical result returned to the LLM.
*/
buildToolHistoryContent(
_args: Record<string, unknown> | null,
_toolResult: ToolResult,
): string | undefined {
return undefined;
}
/**
* Optionally build a user-facing confirmation summary before execution.
* Non-read-only tools should override this with a concise approval prompt.
@@ -0,0 +1,102 @@
import { beforeEach, describe, expect, it, vi } from 'vitest';
import { GetUserSelectedMusicRangeAndTrackTool } from './GetUserSelectedMusicRangeAndTrackTool';
import { KGProject } from '../../core/KGProject';
import { KGMidiTrack } from '../../core/track/KGMidiTrack';
import { KGMidiRegion } from '../../core/region/KGMidiRegion';
import { KGGlobalTrack, GlobalTrackType } from '../../core/global-track/KGGlobalTrack';
import { KGMarkerRegion } from '../../core/region/KGMarkerRegion';
import { KGCore } from '../../core/KGCore';
const storeState = {
activeRegionId: null as string | null,
selectedRegionIds: [] as string[],
selectedTrackId: null as string | null,
};
vi.mock('../../stores/projectStore', () => ({
useProjectStore: {
getState: () => storeState,
},
}));
function buildProject(): KGProject {
const project = new KGProject('selection-tool-project', 8, 0, 120, { numerator: 4, denominator: 4 }, 'C major');
const midiTrack = new KGMidiTrack('Lead', 1);
midiTrack.setRegions([
new KGMidiRegion('midi-a', midiTrack.getId().toString(), midiTrack.getTrackIndex(), 'A', 4, 8),
new KGMidiRegion('midi-b', midiTrack.getId().toString(), midiTrack.getTrackIndex(), 'B', 20, 4),
]);
project.setTracks([midiTrack]);
const markerTrack = new KGGlobalTrack('global-marker', 0, GlobalTrackType.Marker, 'Marker');
markerTrack.setRegions([
new KGMarkerRegion('global-a', markerTrack.getId(), markerTrack.getTrackIndex(), 'Marker A', 2, 2),
]);
project.setGlobalTracks(project.getGlobalTracks().map(track => (
track.getType() === GlobalTrackType.Marker ? markerTrack : track
)));
return project;
}
function mockCore(project: KGProject) {
vi.spyOn(KGCore, 'instance').mockReturnValue({
getCurrentProject: () => project,
} as unknown as KGCore);
}
describe('GetUserSelectedMusicRangeAndTrackTool', () => {
beforeEach(() => {
storeState.activeRegionId = null;
storeState.selectedRegionIds = [];
storeState.selectedTrackId = null;
vi.restoreAllMocks();
});
it('returns the selected music range and selected regular track', async () => {
const project = buildProject();
mockCore(project);
storeState.selectedRegionIds = ['midi-a', 'midi-b'];
storeState.selectedTrackId = '1';
const tool = new GetUserSelectedMusicRangeAndTrackTool();
const result = await tool.execute({});
expect(result.success).toBe(true);
expect(result.result).toBe(
'Current Selected Music Range:\n- Start Beat: 4\n- End Beat: 24\n\nCurrent Selected Track:\ntrack_id: 1\ntrack_name: Lead',
);
});
it('reports no selected track when the selection is global-only even if selectedTrackId is stale', async () => {
const project = buildProject();
mockCore(project);
storeState.selectedRegionIds = ['global-a'];
storeState.selectedTrackId = '1';
const tool = new GetUserSelectedMusicRangeAndTrackTool();
const result = await tool.execute({});
expect(result.success).toBe(true);
expect(result.result).toBe(
'Current Selected Music Range:\n- Start Beat: 2\n- End Beat: 4\n\nCurrent Selected Track:\nNo selected track.',
);
});
it('uses loop bounds and reports no selected track when nothing is selected', async () => {
const project = buildProject();
project.setIsLooping(true);
project.setLoopingRange([2, 5]);
mockCore(project);
const tool = new GetUserSelectedMusicRangeAndTrackTool();
const result = await tool.execute({});
expect(result.success).toBe(true);
expect(result.result).toBe(
'Current Selected Music Range:\n- Start Beat: 8\n- End Beat: 24\n\nCurrent Selected Track:\nNo selected track.',
);
});
});
@@ -0,0 +1,31 @@
import { BaseTool } from './BaseTool';
import type { ToolParameter, ToolResult } from './BaseTool';
import {
resolveSelectedMusicRangeContext,
resolveSelectedTrackContext,
} from './toolTargeting';
export class GetUserSelectedMusicRangeAndTrackTool extends BaseTool {
readonly name = 'get_user_selected_music_range_and_track';
readonly description =
'Get the current selected music range and the current selected regular track, if one is selected. Use this when selection context matters. When you are editing notes on the currently selected track, you do not need to pass track_id or track_name to note-editing tools.';
readonly parameters: Record<string, ToolParameter> = {};
async execute(_params: Record<string, unknown>): Promise<ToolResult> {
try {
const selectedMusicRange = resolveSelectedMusicRangeContext();
const selectedTrack = resolveSelectedTrackContext();
const selectedTrackSection = selectedTrack.hasSelectedTrack
? `track_id: ${selectedTrack.trackId}\ntrack_name: ${selectedTrack.trackName}`
: 'No selected track.';
return this.createSuccessResult(
`Current Selected Music Range:\n${selectedMusicRange.section}\n\nCurrent Selected Track:\n${selectedTrackSection}`,
);
} catch (error) {
return this.createErrorResult(`Failed to read current selected music range and track: ${error}`);
}
}
}
+53
View File
@@ -0,0 +1,53 @@
import { beforeEach, describe, expect, it, vi } from 'vitest';
import { ListAllTracksTool } from './ListAllTracksTool';
import { KGCore } from '../../core/KGCore';
import { KGProject } from '../../core/KGProject';
import { KGMidiTrack } from '../../core/track/KGMidiTrack';
import { KGAudioTrack } from '../../core/track/KGAudioTrack';
function mockCore(project: KGProject) {
vi.spyOn(KGCore, 'instance').mockReturnValue({
getCurrentProject: () => project,
} as unknown as KGCore);
}
describe('ListAllTracksTool', () => {
beforeEach(() => {
vi.restoreAllMocks();
});
it('lists all MIDI tracks with English instrument names', async () => {
const lead = new KGMidiTrack('Lead', 1, 'acoustic_grand_piano');
const bass = new KGMidiTrack('Bass', 2, 'electric_bass_finger');
const audio = new KGAudioTrack('Vocal', 3);
lead.setTrackIndex(0);
bass.setTrackIndex(1);
audio.setTrackIndex(2);
const project = new KGProject('track-list-project');
project.setTracks([lead, bass, audio]);
mockCore(project);
const tool = new ListAllTracksTool();
const result = await tool.execute({});
expect(result.success).toBe(true);
expect(result.result).toBe(
'track_id: 1\ntrack_name: Lead\ninstrument: Acoustic Grand Piano\n\ntrack_id: 2\ntrack_name: Bass\ninstrument: Electric Bass (finger)',
);
});
it('returns a friendly message when there are no MIDI tracks', async () => {
const project = new KGProject('no-midi-tracks');
project.setTracks([new KGAudioTrack('Mixdown', 1)]);
mockCore(project);
const tool = new ListAllTracksTool();
const result = await tool.execute({});
expect(result).toEqual({
success: true,
result: 'No MIDI tracks found.',
});
});
});
+40
View File
@@ -0,0 +1,40 @@
import { BaseTool } from './BaseTool';
import type { ToolParameter, ToolResult } from './BaseTool';
import { KGMidiTrack } from '../../core/track/KGMidiTrack';
import { FLUIDR3_INSTRUMENT_MAP } from '../../constants/generalMidiConstants';
export class ListAllTracksTool extends BaseTool {
readonly name = 'list_all_tracks';
readonly description =
'List all MIDI tracks in the project with their track_id, track_name, and instrument name in English. Use this when you need to inspect available target tracks before choosing one.';
readonly parameters: Record<string, ToolParameter> = {};
async execute(_params: Record<string, unknown>): Promise<ToolResult> {
try {
const midiTracks = this.getCurrentProject().getTracks().filter(
(track): track is KGMidiTrack => track instanceof KGMidiTrack,
);
if (midiTracks.length === 0) {
return this.createSuccessResult('No MIDI tracks found.');
}
const result = midiTracks.map(track => {
const instrumentKey = track.getInstrument();
const instrumentName = FLUIDR3_INSTRUMENT_MAP[instrumentKey]?.displayName ?? instrumentKey;
const trackName = track.getName() || `Track ${track.getTrackIndex() + 1}`;
return [
`track_id: ${track.getId().toString()}`,
`track_name: ${trackName}`,
`instrument: ${instrumentName}`,
].join('\n');
}).join('\n\n');
return this.createSuccessResult(result);
} catch (error) {
return this.createErrorResult(`Failed to list tracks: ${error}`);
}
}
}
@@ -30,7 +30,6 @@ function buildProjectWithRegionAndOptionalChords(chords: string[] = []): {
const chordTrack = findGlobalTrackByType(project, GlobalTrackType.Chord);
expect(chordTrack).not.toBeNull();
chords.forEach((symbol, index) => {
chordTrack!.addRegion(new KGChordRegion(`chord-${index}`, chordTrack!.getId(), chordTrack!.getTrackIndex(), symbol, index * 4, 4));
});
@@ -59,27 +58,10 @@ describe('ReadChordProgressionTool', () => {
expect(result.success).toBe(true);
expect(result.result).toContain('Chord-symbol representation:');
expect(result.result).toContain('[Am]4 | [F]4 | [Dm]4 | [E7]4 | [Am]4 | [C]4 | [Dm]4 | [E7]4 |');
expect(result.result).toContain('[A, C E]4 | [F, A, C]4 | [D F A]4 | [E ^G B d]4 | [A, C E]4 | [C E G]4 | [D F A]4 | [E ^G B d]4 |');
});
it('falls back to the selected MIDI region when no active region exists', async () => {
const { project, midiRegion } = buildProjectWithRegionAndOptionalChords(['Am']);
vi.spyOn(KGCore, 'instance').mockReturnValue({
getCurrentProject: () => project,
getSelectedItems: () => [midiRegion],
} as unknown as KGCore);
const tool = new ReadChordProgressionTool();
const result = await tool.execute({});
expect(result.success).toBe(true);
expect(result.result).toContain('[Am]4 |');
});
it('returns guidance when no chord progression is defined for the region range', async () => {
const { project, midiRegion } = buildProjectWithRegionAndOptionalChords();
storeState.activeRegionId = midiRegion.getId();
it('reads the full chord track when no MIDI region is selected', async () => {
const { project } = buildProjectWithRegionAndOptionalChords(['Am', 'F', 'Dm']);
vi.spyOn(KGCore, 'instance').mockReturnValue({
getCurrentProject: () => project,
@@ -90,12 +72,12 @@ describe('ReadChordProgressionTool', () => {
const result = await tool.execute({});
expect(result.success).toBe(true);
expect(result.result).toContain('No chord progression is defined for the selected MIDI region range.');
expect(result.result).toContain('read_music');
expect(result.result).toContain('[Am]4 | [F]4 | [Dm]4 |');
expect(tool.buildToolResultDisplayContent({}, result)).toBe('Read the chord progression from bars 1 to 3.');
});
it('returns a clear error when no active or selected MIDI region exists', async () => {
const { project } = buildProjectWithRegionAndOptionalChords(['Am']);
it('returns no-chord-defined guidance when the chord track has no chord regions', async () => {
const { project } = buildProjectWithRegionAndOptionalChords();
vi.spyOn(KGCore, 'instance').mockReturnValue({
getCurrentProject: () => project,
@@ -105,22 +87,11 @@ describe('ReadChordProgressionTool', () => {
const tool = new ReadChordProgressionTool();
const result = await tool.execute({});
expect(result.success).toBe(false);
expect(result.result).toContain('No active or selected MIDI region found');
});
it('builds a compact summary for the resolved region span', () => {
const { project, midiRegion } = buildProjectWithRegionAndOptionalChords(['Am']);
storeState.activeRegionId = midiRegion.getId();
vi.spyOn(KGCore, 'instance').mockReturnValue({
getCurrentProject: () => project,
getSelectedItems: () => [],
} as unknown as KGCore);
const tool = new ReadChordProgressionTool();
const summary = tool.buildToolResultDisplayContent({}, { success: true, result: 'raw result' });
expect(summary).toBe('Read the chord progression from bars 1 to 8.');
expect(result.success).toBe(true);
expect(result.result).toBe('No chord progression is defined for the requested range on the global chord track. Use read_music to inspect the notes directly.');
expect(tool.buildToolHistoryContent({}, result)).toBe(
'No chord progression is defined for that range on the global chord track. Use read_music to inspect the notes directly.',
);
expect(tool.buildToolResultDisplayContent({}, result)).toBe('No chord progression is defined for that range.');
});
});
+69 -43
View File
@@ -1,79 +1,105 @@
import { BaseTool } from './BaseTool';
import type { ToolResult, ToolParameter } from './BaseTool';
import { KGMidiRegion } from '../../core/region/KGMidiRegion';
import { useProjectStore } from '../../stores/projectStore';
import { KGCore } from '../../core/KGCore';
import { resolveActiveOrSelectedMidiRegionContext } from './toolTargeting';
import { convertBeatRangeChordProgressionToABCNotation } from '../../util/abcNotationUtil';
import { GlobalTrackType } from '../../core/global-track';
import { KGChordRegion } from '../../core/region/KGChordRegion';
import { findGlobalTrackByType } from '../../util/globalTrackUtil';
interface ChordProgressionRange {
startBeat: number;
endBeat: number;
scope: 'region' | 'song';
}
const NO_CHORD_PROGRESSION_RAW_MESSAGE =
'No chord progression is defined for the requested range on the global chord track. Use read_music to inspect the notes directly.';
const NO_CHORD_PROGRESSION_HISTORY_MESSAGE =
'No chord progression is defined for that range on the global chord track. Use read_music to inspect the notes directly.';
const NO_CHORD_PROGRESSION_UI_MESSAGE =
'No chord progression is defined for that range.';
/**
* Tool for reading user-defined chord progression content from the global chord track.
*/
export class ReadChordProgressionTool extends BaseTool {
readonly name = 'read_chord_progression';
readonly description = 'Read the user-defined chord progression for the currently active or selected MIDI region. The output has two representations of the same progression: first symbolic chord names such as Em7b5, then note-based ABC chord tokens. Chord progression data comes only from chord regions the user defined on the global chord track, so it may be empty. If no chord progression is defined for this range, read the notes directly with read_music.';
readonly description = 'Read the user-defined chord progression from the global chord track. If a MIDI region is active or selected, read the progression for that region. Otherwise, read the full song progression from bar 1 through the last chord region on the chord track.';
readonly parameters: Record<string, ToolParameter> = {};
buildToolResultDisplayContent(args: Record<string, unknown> | null, toolResult: ToolResult): string | undefined {
override buildToolResultDisplayContent(args: Record<string, unknown> | null, toolResult: ToolResult): string | undefined {
void args;
if (!toolResult.success) {
return undefined;
}
const targetRegion = this.findTargetRegion();
if (!targetRegion) {
if (toolResult.result === NO_CHORD_PROGRESSION_RAW_MESSAGE) {
return NO_CHORD_PROGRESSION_UI_MESSAGE;
}
const range = this.resolveReadRange();
if (!range) {
return undefined;
}
const beatsPerBar = this.getCurrentProject().getTimeSignature().numerator;
const startBar = Math.floor(targetRegion.getStartFromBeat() / beatsPerBar) + 1;
const endBar = Math.max(1, Math.ceil((targetRegion.getStartFromBeat() + targetRegion.getLength()) / beatsPerBar));
const barRange = startBar === endBar ? `bar ${startBar}` : `bars ${startBar} to ${endBar}`;
void args;
const startBar = Math.floor(range.startBeat / beatsPerBar) + 1;
const endBar = Math.max(1, Math.ceil(range.endBeat / beatsPerBar));
return `Read the chord progression from ${startBar === endBar ? `bar ${startBar}` : `bars ${startBar} to ${endBar}`}.`;
}
return `Read the chord progression from ${barRange}.`;
override buildToolHistoryContent(args: Record<string, unknown> | null, toolResult: ToolResult): string | undefined {
void args;
if (toolResult.result === NO_CHORD_PROGRESSION_RAW_MESSAGE) {
return NO_CHORD_PROGRESSION_HISTORY_MESSAGE;
}
return undefined;
}
async execute(_params: Record<string, unknown>): Promise<ToolResult> {
try {
const targetRegion = this.findTargetRegion();
if (!targetRegion) {
return this.createErrorResult(
'No active or selected MIDI region found. Please open the piano roll with a region or select a MIDI region first.'
);
const range = this.resolveReadRange();
if (!range) {
return this.createSuccessResult(NO_CHORD_PROGRESSION_RAW_MESSAGE);
}
const project = this.getCurrentProject();
const startBeat = targetRegion.getStartFromBeat();
const endBeat = startBeat + targetRegion.getLength();
const result = convertBeatRangeChordProgressionToABCNotation(project, startBeat, endBeat);
const result = convertBeatRangeChordProgressionToABCNotation(
this.getCurrentProject(),
range.startBeat,
range.endBeat,
);
return this.createSuccessResult(result);
} catch (error) {
return this.createErrorResult(`Failed to read chord progression: ${error}`);
}
}
private findTargetRegion(): KGMidiRegion | null {
const project = this.getCurrentProject();
const tracks = project.getTracks();
const storeState = useProjectStore.getState();
if (storeState.activeRegionId) {
for (const track of tracks) {
const region = track.getRegions().find(candidate => candidate.getId() === storeState.activeRegionId);
if (region instanceof KGMidiRegion) {
return region;
}
}
private resolveReadRange(): ChordProgressionRange | null {
const resolvedRegion = resolveActiveOrSelectedMidiRegionContext();
if (resolvedRegion) {
return {
startBeat: resolvedRegion.region.getStartFromBeat(),
endBeat: resolvedRegion.region.getStartFromBeat() + resolvedRegion.region.getLength(),
scope: 'region',
};
}
const selectedItems = KGCore.instance().getSelectedItems();
for (const item of selectedItems) {
if (item instanceof KGMidiRegion) {
return item;
}
const chordTrack = findGlobalTrackByType(this.getCurrentProject(), GlobalTrackType.Chord);
if (!chordTrack) {
return null;
}
return null;
const chordRegions = chordTrack.getRegions().filter((region): region is KGChordRegion => region instanceof KGChordRegion);
if (chordRegions.length === 0) {
return null;
}
return {
startBeat: 0,
endBeat: Math.max(...chordRegions.map(region => region.getStartFromBeat() + region.getLength())),
scope: 'song',
};
}
}
+22 -38
View File
@@ -19,33 +19,11 @@ describe('ReadMusicTool', () => {
vi.restoreAllMocks();
});
it('builds a compact summary for a single track read', () => {
it('builds a compact summary for reading multiple tracks including empty ones', () => {
const project = new KGProject('read-music-project', 8, 0, 120, { numerator: 4, denominator: 4 }, 'C major');
const leadTrack = buildTrack('Lead', 1, 0, 16);
project.setTracks([leadTrack]);
vi.spyOn(KGCore, 'instance').mockReturnValue({
getCurrentProject: () => project,
} as unknown as KGCore);
const tool = new ReadMusicTool();
const summary = tool.buildToolResultDisplayContent(
{
track_id: leadTrack.getId().toString(),
start: 5,
length: 6,
},
{ success: true, result: 'raw result' },
);
expect(summary).toBe('Read track Lead from bars 2 to 3.');
});
it('builds a compact summary for reading multiple tracks', () => {
const project = new KGProject('read-music-project', 8, 0, 120, { numerator: 4, denominator: 4 }, 'C major');
const leadTrack = buildTrack('Lead', 1, 0, 16);
const bassTrack = buildTrack('Bass', 2, 0, 12);
project.setTracks([leadTrack, bassTrack]);
const emptyTrack = new KGMidiTrack('Pads', 2);
project.setTracks([leadTrack, emptyTrack]);
vi.spyOn(KGCore, 'instance').mockReturnValue({
getCurrentProject: () => project,
@@ -61,28 +39,33 @@ describe('ReadMusicTool', () => {
{ success: true, result: 'raw result' },
);
expect(summary).toBe('Read tracks Lead and Bass from bars 1 to 4.');
expect(summary).toBe('Read tracks Lead and Pads from bars 1 to 4.');
});
it('returns no compact summary when the requested track cannot be resolved', () => {
it('includes empty MIDI tracks as rest-only ABC sections in all-track reads', async () => {
const project = new KGProject('read-music-project', 8, 0, 120, { numerator: 4, denominator: 4 }, 'C major');
project.setTracks([buildTrack('Lead', 1, 0, 16)]);
const leadTrack = buildTrack('Lead', 1, 0, 8);
const emptyTrack = new KGMidiTrack('Pads', 2);
project.setTracks([leadTrack, emptyTrack]);
vi.spyOn(KGCore, 'instance').mockReturnValue({
getCurrentProject: () => project,
} as unknown as KGCore);
const tool = new ReadMusicTool();
const summary = tool.buildToolResultDisplayContent(
{
track_id: 'missing-track',
start: 0,
length: 4,
},
{ success: true, result: 'raw result' },
);
const result = await tool.execute({ track_id: 'all', start: 0, length: 8 });
expect(summary).toBeUndefined();
expect(result.success).toBe(true);
expect(result.result).toContain('track_id: 1');
expect(result.result).toContain('track_name: Lead');
expect(result.result).toContain('Instrument: Acoustic Grand Piano');
expect(result.result).toContain('track_id: 2');
expect(result.result).toContain('track_name: Pads');
expect(result.result).toContain('Instrument: Acoustic Grand Piano');
expect(result.result).not.toContain('Track 1 - Melody:');
expect(result.result).not.toContain('Track 2 - Pads:');
expect(result.result).not.toContain('\nT:');
expect(result.result).toContain('z4 | z4 | // No regions found');
});
it('returns a professional empty-project message when all MIDI tracks are empty', async () => {
@@ -107,7 +90,8 @@ describe('ReadMusicTool', () => {
it('returns a professional empty-range message when the selected range has no MIDI notes', async () => {
const project = new KGProject('read-music-project', 8, 0, 120, { numerator: 4, denominator: 4 }, 'C major');
const leadTrack = buildTrack('Lead', 1, 0, 8);
project.setTracks([leadTrack]);
const emptyTrack = new KGMidiTrack('Pads', 2);
project.setTracks([leadTrack, emptyTrack]);
vi.spyOn(KGCore, 'instance').mockReturnValue({
getCurrentProject: () => project,
+50 -119
View File
@@ -179,10 +179,7 @@ export class ReadMusicTool extends BaseTool {
if (!trackId || trackId === '' || trackId === 'all') {
const midiTracks = tracks.filter(track => track instanceof KGMidiTrack) as KGMidiTrack[];
const tracksToSkip = this.findTracksToSkip(midiTracks);
return midiTracks
.filter(track => !tracksToSkip.includes(track))
.map((track, index) => track.getName() || `Track ${index + 1}`);
return midiTracks.map((track, index) => track.getName() || `Track ${index + 1}`);
}
const targetTrack = tracks.find(track => track.getId().toString() === trackId);
@@ -202,9 +199,7 @@ export class ReadMusicTool extends BaseTool {
if (!trackId || trackId === '' || trackId === 'all') {
const midiTracks = tracks.filter(track => track instanceof KGMidiTrack) as KGMidiTrack[];
const tracksToSkip = this.findTracksToSkip(midiTracks);
const visibleTracks = midiTracks.filter(track => !tracksToSkip.includes(track));
const endBeats = visibleTracks.flatMap(track =>
const endBeats = midiTracks.flatMap(track =>
track.getRegions()
.filter(region => region instanceof KGMidiRegion)
.map(region => region.getStartFromBeat() + region.getLength())
@@ -264,44 +259,41 @@ export class ReadMusicTool extends BaseTool {
return 'No musical content was found in the selected range.';
}
/**
* Find tracks that should be skipped because they have no musical content
* (no regions or regions with no notes)
*/
private findTracksToSkip(tracks: KGMidiTrack[]): KGMidiTrack[] {
try {
const tracksToSkip: KGMidiTrack[] = [];
private buildTrackHeader(track: KGMidiTrack): string {
const project = this.getCurrentProject();
const timeSignature = project.getTimeSignature();
const bpm = project.getBpm();
const keySignature = project.getKeySignature();
const abcKeySignature = KEY_SIGNATURE_MAP[keySignature]?.abcNotationKeySignature || 'C';
const trackId = track.getId().toString();
const trackName = track.getName() || 'Unnamed Track';
const instrumentName = FLUIDR3_INSTRUMENT_MAP[track.getInstrument()]?.displayName || track.getInstrument();
for (const track of tracks) {
const regions = track.getRegions();
// Skip tracks with no regions
if (regions.length === 0) {
tracksToSkip.push(track);
continue;
}
// Check if all regions in this track are empty (have no notes)
const hasAnyNotes = regions.some(region => {
if (region.getCurrentType() === 'KGMidiRegion') {
return (region as KGMidiRegion).getNotes().length > 0;
}
return false;
});
// Skip tracks where no regions have notes
if (!hasAnyNotes) {
tracksToSkip.push(track);
}
}
return tracksToSkip;
} catch (error) {
console.error('Error finding tracks to skip:', error);
return [];
}
return [
`track_id: ${trackId}`,
`track_name: ${trackName}`,
`Instrument: ${instrumentName}`,
'X:1',
`M:${timeSignature.numerator}/${timeSignature.denominator}`,
`L:1/${timeSignature.denominator}`,
`Q:1/${timeSignature.denominator}=${bpm}`,
`K:${abcKeySignature}`
].join('\n');
}
private hasAnyMidiNotes(tracks: KGMidiTrack[]): boolean {
return tracks.some(track => track.getRegions().some(region => (
region instanceof KGMidiRegion && region.getNotes().length > 0
)));
}
private buildRestBody(startBeat: number, endBeat: number | undefined): string {
const beatsPerBar = this.getCurrentProject().getTimeSignature().numerator;
const effectiveEndBeat = endBeat ?? (startBeat + beatsPerBar);
const totalBars = Math.max(1, Math.ceil((effectiveEndBeat - startBeat) / beatsPerBar));
const restToken = `z${beatsPerBar}`;
return Array.from({ length: totalBars }, () => restToken).join(' | ') + ' |';
}
/**
* Generate ABC notation for all tracks
@@ -313,59 +305,24 @@ export class ReadMusicTool extends BaseTool {
return this.getEmptyProjectMessage();
}
// Find tracks to skip (tracks with no content)
const tracksToSkip = this.findTracksToSkip(midiTracks);
const visibleTracks = midiTracks.filter(track => !tracksToSkip.includes(track));
if (visibleTracks.length === 0) {
if (!this.hasAnyMidiNotes(midiTracks)) {
return this.getEmptyProjectMessage();
}
const hasContentInRange = visibleTracks.some(track => this.hasMidiContentInRange(track, startBeat, endBeat));
const hasContentInRange = midiTracks.some(track => this.hasMidiContentInRange(track, startBeat, endBeat));
if (!hasContentInRange) {
return this.getEmptyRangeMessage();
}
// Get project settings for proper notation
const project = this.getCurrentProject();
const timeSignature = project.getTimeSignature();
const keySignature = project.getKeySignature();
const abcKeySignature = KEY_SIGNATURE_MAP[keySignature]?.abcNotationKeySignature || 'C';
let output = `All Tracks (beats ${startBeat}-${endBeat || 'end'}):\n\n`;
midiTracks.forEach((track, index) => {
// Skip tracks that have no musical content
if (tracksToSkip.includes(track)) {
return; // Skip this track
}
const trackNumber = index + 1;
const trackName = track.getName() || `Track ${trackNumber}`;
// Check if this track uses a percussion instrument
const percussionDisplayName = this.getPercussionDisplayName(track);
let displayTrackName: string;
if (percussionDisplayName) {
// Use percussion instrument display name for all percussion tracks
displayTrackName = percussionDisplayName;
} else if (trackNumber === 1) {
// Use "Melody" for the first non-percussion track
displayTrackName = 'Melody';
} else {
// Use original track name for other non-percussion tracks
displayTrackName = trackName;
}
output += `Track ${trackNumber} - ${displayTrackName}:\n`;
let output = `Tracks (beats ${startBeat}-${endBeat || 'end'}):\n\n`;
midiTracks.forEach((track) => {
// Get all regions from the track and convert each one
const regions = track.getRegions().filter(region => region instanceof KGMidiRegion) as KGMidiRegion[];
if (regions.length === 0) {
output += 'X:' + trackNumber + '\n';
output += `M:${timeSignature.numerator}/${timeSignature.denominator}\n`;
output += `K:${abcKeySignature}\n`;
output += 'z4 | // No regions found\n\n';
output += `${this.buildTrackHeader(track)}\n`;
output += `${this.buildRestBody(startBeat, endBeat)} // No regions found\n\n`;
} else {
// Convert each region that overlaps with the requested range
let hasContent = false;
@@ -376,20 +333,14 @@ export class ReadMusicTool extends BaseTool {
// Check if region overlaps with requested range
if (regionStart < (endBeat || Infinity) && regionEnd > startBeat) {
const abcNotation = convertRegionToABCNotation(region, startBeat, endBeat);
// Update the X: line to include track number
const lines = abcNotation.split('\n');
lines[0] = `X:${trackNumber}`;
output += lines.join('\n') + '\n\n';
output += abcNotation + '\n\n';
hasContent = true;
}
});
if (!hasContent) {
output += 'X:' + trackNumber + '\n';
output += `M:${timeSignature.numerator}/${timeSignature.denominator}\n`;
output += `K:${abcKeySignature}\n`;
output += 'z4 | // No content in specified range\n\n';
output += `${this.buildTrackHeader(track)}\n`;
output += `${this.buildRestBody(startBeat, endBeat)} // No content in specified range\n\n`;
}
}
});
@@ -405,7 +356,7 @@ export class ReadMusicTool extends BaseTool {
return `Track is not a MIDI track.`;
}
if (track.getRegions().length === 0 || !track.getRegions().some(region => (
if (!track.getRegions().some(region => (
region instanceof KGMidiRegion && region.getNotes().length > 0
))) {
return this.getEmptyProjectMessage();
@@ -415,26 +366,14 @@ export class ReadMusicTool extends BaseTool {
return this.getEmptyRangeMessage();
}
// Get project settings for proper notation
const project = this.getCurrentProject();
const timeSignature = project.getTimeSignature();
const keySignature = project.getKeySignature();
const abcKeySignature = KEY_SIGNATURE_MAP[keySignature]?.abcNotationKeySignature || 'C';
const trackName = track.getName() || 'Unnamed Track';
let output = `Track "${trackName}" (beats ${startBeat}-${endBeat || 'end'}):\n`;
let output = '';
// Get all regions from the track and convert each one
const regions = track.getRegions().filter(region => region instanceof KGMidiRegion) as KGMidiRegion[];
if (regions.length === 0) {
output += 'X:1\n';
output += `T:${trackName}\n`;
output += `M:${timeSignature.numerator}/${timeSignature.denominator}\n`;
output += `K:${abcKeySignature}\n`;
output += `L:1/${timeSignature.denominator}\n`;
output += 'z4 | // No regions found';
output += `${this.buildTrackHeader(track)}\n`;
output += `${this.buildRestBody(startBeat, endBeat)} // No regions found`;
} else {
// Convert each region that overlaps with the requested range
let hasContent = false;
@@ -445,22 +384,14 @@ export class ReadMusicTool extends BaseTool {
// Check if region overlaps with requested range
if (regionStart < (endBeat || Infinity) && regionEnd > startBeat) {
const abcNotation = convertRegionToABCNotation(region, startBeat, endBeat);
// Update the title to include track name
const lines = abcNotation.split('\n');
lines[1] = `T:${trackName}`;
output += lines.join('\n');
output += abcNotation;
hasContent = true;
}
});
if (!hasContent) {
output += 'X:1\n';
output += `T:${trackName}\n`;
output += `M:${timeSignature.numerator}/${timeSignature.denominator}\n`;
output += `K:${abcKeySignature}\n`;
output += `L:1/${timeSignature.denominator}\n`;
output += 'z4 | // No content in specified range';
output += `${this.buildTrackHeader(track)}\n`;
output += `${this.buildRestBody(startBeat, endBeat)} // No content in specified range`;
}
}
+130 -13
View File
@@ -1,5 +1,10 @@
import { beforeEach, describe, expect, it, vi } from 'vitest';
import { RemoveNotesTool } from './RemoveNotesTool';
import {
NO_MIDI_TARGET_HISTORY_MESSAGE,
NO_MIDI_TARGET_RAW_MESSAGE,
NO_MIDI_TARGET_UI_MESSAGE,
} from './toolTargeting';
import { KGCore } from '../../core/KGCore';
import { KGProject } from '../../core/KGProject';
import { KGMidiTrack } from '../../core/track/KGMidiTrack';
@@ -16,13 +21,21 @@ vi.mock('../../stores/projectStore', () => ({
},
}));
function mockCore(project: KGProject, selectedItems: unknown[] = []) {
vi.spyOn(KGCore, 'instance').mockReturnValue({
getCurrentProject: () => project,
getSelectedItems: () => selectedItems,
executeCommand: (command: { execute(): void }) => command.execute(),
} as unknown as KGCore);
}
describe('RemoveNotesTool', () => {
beforeEach(() => {
storeState.activeRegionId = null;
vi.restoreAllMocks();
});
it('builds confirmation and result summaries for note removal', () => {
it('builds confirmation and result summaries for region-scoped removal', () => {
const track = new KGMidiTrack('Lead', 1);
const region = new KGMidiRegion('region-1', track.getId().toString(), track.getTrackIndex(), 'Verse Melody', 0, 32);
region.setNotes([
@@ -33,24 +46,128 @@ describe('RemoveNotesTool', () => {
const project = new KGProject('summary-project', 8, 0, 120, { numerator: 4, denominator: 4 }, 'C major');
project.setTracks([track]);
storeState.activeRegionId = region.getId();
vi.spyOn(KGCore, 'instance').mockReturnValue({
getCurrentProject: () => project,
getSelectedItems: () => [],
} as unknown as KGCore);
mockCore(project);
const tool = new RemoveNotesTool();
const args = {
start: 16,
end: 24,
};
const args = { start: 16, end: 24 };
expect(tool.isReadOnlyTool()).toBe(false);
expect(tool.buildConfirmationContent(args)).toBe(
'Allow removing 2 notes from beats 16-24, in region **Verse Melody** on track **Lead**, spanning bars 5 to 7?'
'Allow removing 2 notes from beats 16-24, in region **Verse Melody** on track **Lead**, spanning bars 5 to 7?',
);
expect(tool.buildToolResultDisplayContent(args, { success: true, result: 'raw result' })).toBe(
'Successfully removed 2 notes from beats 16-24, in region **Verse Melody** on track **Lead**, spanning bars 5 to 7.'
'Successfully removed 2 notes from beats 16-24, in region **Verse Melody** on track **Lead**, spanning bars 5 to 7.',
);
});
it('removes notes across all MIDI regions on a track when track_id is provided', async () => {
const track = new KGMidiTrack('Lead', 1);
const regionA = new KGMidiRegion('region-a', track.getId().toString(), track.getTrackIndex(), 'A', 0, 8);
const regionB = new KGMidiRegion('region-b', track.getId().toString(), track.getTrackIndex(), 'B', 8, 8);
regionA.setNotes([new KGMidiNote('note-1', 2, 3, 60, 100)]);
regionB.setNotes([new KGMidiNote('note-2', 2, 3, 64, 100)]);
track.setRegions([regionA, regionB]);
const project = new KGProject('track-remove-project', 8, 0, 120, { numerator: 4, denominator: 4 }, 'C major');
project.setTracks([track]);
mockCore(project);
const tool = new RemoveNotesTool();
const result = await tool.execute({
track_id: track.getId().toString(),
start: 0,
end: 12,
});
expect(result.success).toBe(true);
expect(regionA.getNotes()).toHaveLength(0);
expect(regionB.getNotes()).toHaveLength(0);
});
it('removes notes across a track resolved by track_name when track_id is omitted', async () => {
const leadTrack = new KGMidiTrack('Lead', 1);
const bassTrack = new KGMidiTrack('Bass', 2);
const leadRegion = new KGMidiRegion('lead-region', leadTrack.getId().toString(), leadTrack.getTrackIndex(), 'Lead Region', 0, 8);
const bassRegion = new KGMidiRegion('bass-region', bassTrack.getId().toString(), bassTrack.getTrackIndex(), 'Bass Region', 0, 8);
leadRegion.setNotes([new KGMidiNote('lead-note', 1, 2, 60, 100)]);
bassRegion.setNotes([new KGMidiNote('bass-note', 1, 2, 48, 100)]);
leadTrack.setRegions([leadRegion]);
bassTrack.setRegions([bassRegion]);
const project = new KGProject('remove-track-name-project', 8, 0, 120, { numerator: 4, denominator: 4 }, 'C major');
project.setTracks([leadTrack, bassTrack]);
mockCore(project);
const tool = new RemoveNotesTool();
const result = await tool.execute({
track_name: 'Bass',
start: 0,
end: 4,
});
expect(result.success).toBe(true);
expect(leadRegion.getNotes()).toHaveLength(1);
expect(bassRegion.getNotes()).toHaveLength(0);
});
it('uses track_id when both track_id and track_name are provided', async () => {
const leadTrack = new KGMidiTrack('Lead', 1);
const bassTrack = new KGMidiTrack('Bass', 2);
const leadRegion = new KGMidiRegion('lead-region', leadTrack.getId().toString(), leadTrack.getTrackIndex(), 'Lead Region', 0, 8);
const bassRegion = new KGMidiRegion('bass-region', bassTrack.getId().toString(), bassTrack.getTrackIndex(), 'Bass Region', 0, 8);
leadRegion.setNotes([new KGMidiNote('lead-note', 1, 2, 60, 100)]);
bassRegion.setNotes([new KGMidiNote('bass-note', 1, 2, 48, 100)]);
leadTrack.setRegions([leadRegion]);
bassTrack.setRegions([bassRegion]);
const project = new KGProject('remove-track-id-precedence-project', 8, 0, 120, { numerator: 4, denominator: 4 }, 'C major');
project.setTracks([leadTrack, bassTrack]);
mockCore(project);
const tool = new RemoveNotesTool();
const result = await tool.execute({
track_id: leadTrack.getId().toString(),
track_name: 'Bass',
start: 0,
end: 4,
});
expect(result.success).toBe(true);
expect(leadRegion.getNotes()).toHaveLength(0);
expect(bassRegion.getNotes()).toHaveLength(1);
});
it('uses the first matching track when duplicate track names exist', async () => {
const firstLead = new KGMidiTrack('Lead', 1);
const secondLead = new KGMidiTrack('Lead', 2);
const firstRegion = new KGMidiRegion('first-region', firstLead.getId().toString(), firstLead.getTrackIndex(), 'First Lead Region', 0, 8);
const secondRegion = new KGMidiRegion('second-region', secondLead.getId().toString(), secondLead.getTrackIndex(), 'Second Lead Region', 0, 8);
firstRegion.setNotes([new KGMidiNote('first-note', 1, 2, 60, 100)]);
secondRegion.setNotes([new KGMidiNote('second-note', 1, 2, 64, 100)]);
firstLead.setRegions([firstRegion]);
secondLead.setRegions([secondRegion]);
const project = new KGProject('remove-duplicate-track-name-project', 8, 0, 120, { numerator: 4, denominator: 4 }, 'C major');
project.setTracks([firstLead, secondLead]);
mockCore(project);
const tool = new RemoveNotesTool();
const result = await tool.execute({
track_name: 'Lead',
start: 0,
end: 4,
});
expect(result.success).toBe(true);
expect(firstRegion.getNotes()).toHaveLength(0);
expect(secondRegion.getNotes()).toHaveLength(1);
});
it('returns distinct raw, history, and UI guidance when no MIDI target is available', async () => {
const project = new KGProject('no-target-project', 8, 0, 120, { numerator: 4, denominator: 4 }, 'C major');
mockCore(project);
const tool = new RemoveNotesTool();
const args = { start: 0, end: 4 };
const result = await tool.execute(args);
expect(result).toEqual({ success: false, result: NO_MIDI_TARGET_RAW_MESSAGE });
expect(tool.buildToolHistoryContent(args, result)).toBe(NO_MIDI_TARGET_HISTORY_MESSAGE);
expect(tool.buildToolResultDisplayContent(args, result)).toBe(NO_MIDI_TARGET_UI_MESSAGE);
});
});
+188 -154
View File
@@ -1,27 +1,37 @@
import { BaseTool } from './BaseTool';
import type { ToolResult, ToolParameter } from './BaseTool';
import {
NO_MIDI_TARGET_HISTORY_MESSAGE,
NO_MIDI_TARGET_RAW_MESSAGE,
NO_MIDI_TARGET_UI_MESSAGE,
getTrackDisplayName,
resolveMidiTrackByIdOrName,
resolveActiveOrSelectedMidiRegionContext,
} from './toolTargeting';
import { DeleteNotesCommand } from '../../core/commands/note/DeleteNotesCommand';
import { KGMidiNote } from '../../core/midi/KGMidiNote';
import { KGMidiRegion } from '../../core/region/KGMidiRegion';
import { useProjectStore } from '../../stores/projectStore';
import { KGCore } from '../../core/KGCore';
import { KGMidiTrack } from '../../core/track/KGMidiTrack';
interface RemoveTargetRegionContext {
region: KGMidiRegion;
trackName: string;
}
interface RemoveNotesSummaryData {
noteCount: number;
startBeat: number;
endBeat: number;
regionName: string;
regionName?: string;
trackName: string;
earliestNoteStartBar: number;
latestNoteEndBar: number;
scope: 'region' | 'track';
}
/**
* Tool for removing notes from MIDI regions within a specified beat range
* Integrates with the existing command system for undo/redo support
*/
export class RemoveNotesTool extends BaseTool {
readonly name = 'remove_notes';
readonly description = 'Remove all MIDI notes whose start position falls within the specified beat range. Use this to clear a section before rewriting it, or to delete unwanted notes. Beat positions are absolute on the project timeline.';
readonly description = 'Remove MIDI notes from an absolute beat range. Use track_id to remove notes across every MIDI region on a track. If track_id is omitted, the currently active or selected MIDI region is used.';
override isReadOnlyTool(): boolean {
return false;
@@ -31,31 +41,51 @@ export class RemoveNotesTool extends BaseTool {
start: {
type: 'number',
description: 'Start beat — the absolute beat position where the removal range begins (inclusive). A note starting at exactly this beat will be removed.',
required: true
required: true,
},
end: {
type: 'number',
description: 'End beat — the absolute beat position where the removal range ends (exclusive). A note starting at exactly this beat will NOT be removed. Must be greater than start.',
required: true
required: true,
},
region_id: {
track_id: {
type: 'string',
description: 'Target region ID. If omitted, uses the currently active piano roll region or selected region.',
required: false
}
description: 'Optional target MIDI track ID. If provided, matching notes are removed across all MIDI regions on that track whose absolute start positions fall within the requested beat range.',
required: false,
},
track_name: {
type: 'string',
description: 'Optional target MIDI track name. Used only when track_id is omitted. If multiple MIDI tracks share the same name, the first matching track is used.',
required: false,
},
};
override buildToolResultDisplayContent(args: Record<string, unknown> | null, toolResult: ToolResult): string | undefined {
if (!toolResult.success || !args) {
if (!args) {
return undefined;
}
if (!toolResult.success) {
return toolResult.result === NO_MIDI_TARGET_RAW_MESSAGE ? NO_MIDI_TARGET_UI_MESSAGE : undefined;
}
const summary = this.buildSummaryData(args);
if (!summary || summary.noteCount === 0) {
return undefined;
}
return `Successfully removed ${summary.noteCount} ${summary.noteCount === 1 ? 'note' : 'notes'} from beats ${summary.startBeat}-${summary.endBeat}, in region **${summary.regionName}** on track **${summary.trackName}**, spanning bars ${summary.earliestNoteStartBar} to ${summary.latestNoteEndBar}.`;
const location = summary.scope === 'track'
? `on track **${summary.trackName}**`
: `in region **${summary.regionName}** on track **${summary.trackName}**`;
return `Successfully removed ${summary.noteCount} ${summary.noteCount === 1 ? 'note' : 'notes'} from beats ${summary.startBeat}-${summary.endBeat}, ${location}, spanning bars ${summary.earliestNoteStartBar} to ${summary.latestNoteEndBar}.`;
}
override buildToolHistoryContent(args: Record<string, unknown> | null, toolResult: ToolResult): string | undefined {
if (!args || toolResult.success) {
return undefined;
}
return toolResult.result === NO_MIDI_TARGET_RAW_MESSAGE ? NO_MIDI_TARGET_HISTORY_MESSAGE : undefined;
}
override buildConfirmationContent(args: Record<string, unknown> | null): string | undefined {
@@ -68,196 +98,200 @@ export class RemoveNotesTool extends BaseTool {
return undefined;
}
return `Allow removing ${summary.noteCount} ${summary.noteCount === 1 ? 'note' : 'notes'} from beats ${summary.startBeat}-${summary.endBeat}, in region **${summary.regionName}** on track **${summary.trackName}**, spanning bars ${summary.earliestNoteStartBar} to ${summary.latestNoteEndBar}?`;
const location = summary.scope === 'track'
? `on track **${summary.trackName}**`
: `in region **${summary.regionName}** on track **${summary.trackName}**`;
return `Allow removing ${summary.noteCount} ${summary.noteCount === 1 ? 'note' : 'notes'} from beats ${summary.startBeat}-${summary.endBeat}, ${location}, spanning bars ${summary.earliestNoteStartBar} to ${summary.latestNoteEndBar}?`;
}
async execute(params: Record<string, unknown>): Promise<ToolResult> {
try {
// Validate parameters
this.validateParameters(params);
const startBeat = params.start as number;
const endBeat = params.end as number;
const regionId = params.region_id as string | undefined;
// Validate beat range
const trackId = params.track_id as string | undefined;
const trackName = params.track_name as string | undefined;
if (startBeat < 0) {
return this.createErrorResult(`Invalid start ${startBeat}. Must be >= 0.`);
}
if (endBeat <= startBeat) {
return this.createErrorResult(`Invalid beat range: end (${endBeat}) must be greater than start (${startBeat}).`);
}
// Find the target region
const targetRegion = this.findTargetRegion(regionId);
if (!targetRegion) {
return this.createErrorResult(
regionId
? `Region with ID "${regionId}" not found or is not a MIDI region`
: 'No active or selected MIDI region found. Please open the piano roll with a region or select a MIDI region first.'
);
if (trackId || trackName) {
const explicitTrack = resolveMidiTrackByIdOrName(trackId, trackName);
if (!explicitTrack) {
return this.createErrorResult(
trackId
? `Track with ID "${trackId}" not found or is not a MIDI track.`
: `Track with name "${trackName}" not found or is not a MIDI track.`,
);
}
}
// Adjust beat range relative to region's start beat
const regionStartBeat = targetRegion.getStartFromBeat();
const adjustedStartBeat = startBeat - regionStartBeat;
const adjustedEndBeat = endBeat - regionStartBeat;
// Find all notes within the specified beat range
const notesToRemove = this.findNotesInRange(targetRegion, adjustedStartBeat, adjustedEndBeat);
if (notesToRemove.length === 0) {
return this.createSuccessResult(
`No notes found in the range from beat ${startBeat} to ${endBeat}.`
);
const notesToRemove = (trackId || trackName)
? this.findTrackNotesInRange(trackId, trackName, startBeat, endBeat)
: this.findFallbackRegionNotesInRange(startBeat, endBeat);
if (!notesToRemove) {
return this.createErrorResult(NO_MIDI_TARGET_RAW_MESSAGE);
}
// Extract note IDs for deletion
const noteIds = notesToRemove.map(note => note.getId());
// Execute the deletion command
const command = new DeleteNotesCommand(noteIds);
await this.executeCommand(command);
// Create success message
const noteCount = notesToRemove.length;
const noteList = notesToRemove
.map(note => {
const noteNames = ['C', 'C#', 'D', 'D#', 'E', 'F', 'F#', 'G', 'G#', 'A', 'A#', 'B'];
const octave = Math.floor(note.getPitch() / 12) - 1;
const noteName = noteNames[note.getPitch() % 12];
return `${noteName}${octave}`;
})
.join(', ');
if (notesToRemove.notes.length === 0) {
return this.createSuccessResult(`No notes found in the range from beat ${startBeat} to ${endBeat}.`);
}
const noteIds = notesToRemove.notes.map(note => note.getId());
await this.executeCommand(new DeleteNotesCommand(noteIds));
const noteList = notesToRemove.notes.map(note => this.formatMidiPitch(note.getPitch())).join(', ');
const scopeLabel = notesToRemove.scope === 'track'
? `track "${notesToRemove.trackName}"`
: `MIDI region "${notesToRemove.regionName}" on track "${notesToRemove.trackName}"`;
return this.createSuccessResult(
`Successfully removed ${noteCount} note${noteCount > 1 ? 's' : ''} from beats ${startBeat}-${endBeat}: ${noteList}`
`Successfully removed ${notesToRemove.notes.length} note${notesToRemove.notes.length > 1 ? 's' : ''} from beats ${startBeat}-${endBeat} in ${scopeLabel}: ${noteList}`,
);
} catch (error) {
return this.createErrorResult(`Failed to remove notes: ${error}`);
}
}
/**
* Find the target region for note removal
* Priority: 1) Specified regionId, 2) Active piano roll region, 3) Selected regions, 4) Error if none found
*/
private findTargetRegion(regionId?: string): KGMidiRegion | null {
return this.findTargetRegionContext(regionId)?.region ?? null;
}
private findTargetRegionContext(regionId?: string): { region: KGMidiRegion; trackName: string } | null {
const project = this.getCurrentProject();
const tracks = project.getTracks();
if (regionId) {
// Find specific region by ID
for (const track of tracks) {
const regions = track.getRegions();
const region = regions.find(r => r.getId() === regionId);
if (region && region instanceof KGMidiRegion) {
return {
region,
trackName: track.getName() || `Track ${track.getTrackIndex() + 1}`,
};
}
}
return null;
} else {
// Smart region finding: try different sources in priority order
// 1. Try active piano roll region
const storeState = useProjectStore.getState();
if (storeState.activeRegionId) {
for (const track of tracks) {
const regions = track.getRegions();
const region = regions.find(r => r.getId() === storeState.activeRegionId);
if (region && region instanceof KGMidiRegion) {
return {
region,
trackName: track.getName() || `Track ${track.getTrackIndex() + 1}`,
};
}
}
}
// 2. Try selected regions
const core = this.getKGCore();
const selectedItems = core.getSelectedItems();
for (const item of selectedItems) {
if (item instanceof KGMidiRegion) {
const track = tracks.find(candidate => candidate.getId().toString() === item.getTrackId());
return {
region: item,
trackName: track?.getName() || `Track ${item.getTrackIndex() + 1}`,
};
}
}
// 3. No fallback - return null to trigger error
return null;
}
}
private buildSummaryData(args: Record<string, unknown>): RemoveNotesSummaryData | null {
const typedArgs = args as {
start?: number;
end?: number;
region_id?: string;
track_id?: string;
track_name?: string;
};
if (typeof typedArgs.start !== 'number' || typeof typedArgs.end !== 'number' || typedArgs.end <= typedArgs.start) {
return null;
}
const targetRegion = this.findTargetRegionContext(typedArgs.region_id);
if (!targetRegion) {
const notesInRange = (typedArgs.track_id || typedArgs.track_name)
? this.findTrackNotesInRange(typedArgs.track_id, typedArgs.track_name, typedArgs.start, typedArgs.end)
: this.findFallbackRegionNotesInRange(typedArgs.start, typedArgs.end);
if (!notesInRange) {
return null;
}
const regionStartBeat = targetRegion.region.getStartFromBeat();
const adjustedStartBeat = typedArgs.start - regionStartBeat;
const adjustedEndBeat = typedArgs.end - regionStartBeat;
const notesToRemove = this.findNotesInRange(targetRegion.region, adjustedStartBeat, adjustedEndBeat);
const beatsPerBar = this.getCurrentProject().getTimeSignature().numerator;
let earliestBeat = typedArgs.start;
let latestBeat = typedArgs.end;
if (notesToRemove.length > 0) {
earliestBeat = Math.min(...notesToRemove.map(note => note.getStartBeat() + regionStartBeat));
latestBeat = Math.max(...notesToRemove.map(note => note.getEndBeat() + regionStartBeat));
if (notesInRange.notes.length > 0) {
earliestBeat = Math.min(...notesInRange.notes.map(note => notesInRange.absoluteBoundsByNoteId.get(note.getId())!.startBeat));
latestBeat = Math.max(...notesInRange.notes.map(note => notesInRange.absoluteBoundsByNoteId.get(note.getId())!.endBeat));
}
return {
noteCount: notesToRemove.length,
noteCount: notesInRange.notes.length,
startBeat: typedArgs.start,
endBeat: typedArgs.end,
regionName: targetRegion.region.getName(),
trackName: targetRegion.trackName,
regionName: notesInRange.scope === 'region' ? notesInRange.regionName : undefined,
trackName: notesInRange.trackName,
earliestNoteStartBar: Math.floor(earliestBeat / beatsPerBar) + 1,
latestNoteEndBar: Math.max(1, Math.ceil(latestBeat / beatsPerBar)),
scope: notesInRange.scope,
};
}
/**
* Get KGCore instance for selection access
*/
private getKGCore() {
return KGCore.instance();
private findFallbackRegionNotesInRange(startBeat: number, endBeat: number): {
scope: 'region';
trackName: string;
regionName: string;
notes: KGMidiNote[];
absoluteBoundsByNoteId: Map<string, { startBeat: number; endBeat: number }>;
} | null {
const resolvedRegion = this.resolveFallbackRegion();
if (!resolvedRegion) {
return null;
}
const regionStartBeat = resolvedRegion.region.getStartFromBeat();
const adjustedStartBeat = startBeat - regionStartBeat;
const adjustedEndBeat = endBeat - regionStartBeat;
const notes = resolvedRegion.region.getNotes().filter(note => {
const noteStartBeat = note.getStartBeat();
return noteStartBeat >= adjustedStartBeat && noteStartBeat < adjustedEndBeat;
});
return {
scope: 'region',
trackName: resolvedRegion.trackName,
regionName: resolvedRegion.region.getName(),
notes,
absoluteBoundsByNoteId: new Map(notes.map(note => ([
note.getId(),
{
startBeat: note.getStartBeat() + regionStartBeat,
endBeat: note.getEndBeat() + regionStartBeat,
},
]))),
};
}
/**
* Find all notes within the specified beat range
* Notes are included if their start beat is within [startBeat, endBeat)
*/
private findNotesInRange(region: KGMidiRegion, startBeat: number, endBeat: number) {
const notes = region.getNotes();
return notes.filter(note => {
const noteStartBeat = note.getStartBeat();
return noteStartBeat >= startBeat && noteStartBeat < endBeat;
});
private findTrackNotesInRange(trackId: string | undefined, trackName: string | undefined, startBeat: number, endBeat: number): {
scope: 'track';
trackName: string;
notes: KGMidiNote[];
absoluteBoundsByNoteId: Map<string, { startBeat: number; endBeat: number }>;
} | null {
const track = resolveMidiTrackByIdOrName(trackId, trackName);
if (!track) {
return null;
}
const notes: KGMidiNote[] = [];
const absoluteBoundsByNoteId = new Map<string, { startBeat: number; endBeat: number }>();
for (const region of track.getRegions()) {
if (!(region instanceof KGMidiRegion)) {
continue;
}
const regionStartBeat = region.getStartFromBeat();
for (const note of region.getNotes()) {
const absoluteStartBeat = regionStartBeat + note.getStartBeat();
if (absoluteStartBeat < startBeat || absoluteStartBeat >= endBeat) {
continue;
}
notes.push(note);
absoluteBoundsByNoteId.set(note.getId(), {
startBeat: absoluteStartBeat,
endBeat: regionStartBeat + note.getEndBeat(),
});
}
}
return {
scope: 'track',
trackName: getTrackDisplayName(track),
notes,
absoluteBoundsByNoteId,
};
}
private resolveFallbackRegion(): RemoveTargetRegionContext | null {
const resolved = resolveActiveOrSelectedMidiRegionContext();
if (!resolved) {
return null;
}
return {
region: resolved.region,
trackName: resolved.trackName,
};
}
private formatMidiPitch(midiPitch: number): string {
const noteNames = ['C', 'C#', 'D', 'D#', 'E', 'F', 'F#', 'G', 'G#', 'A', 'A#', 'B'];
const octave = Math.floor(midiPitch / 12) - 1;
const noteName = noteNames[midiPitch % 12];
return `${noteName}${octave}`;
}
}
+13 -1
View File
@@ -9,8 +9,18 @@ import { RemoveNotesTool } from './RemoveNotesTool';
import { ReadMusicTool } from './ReadMusicTool';
import { ReadChordProgressionTool } from './ReadChordProgressionTool';
import { UpdateTodoListTool } from './UpdateTodoListTool';
import { GetUserSelectedMusicRangeAndTrackTool } from './GetUserSelectedMusicRangeAndTrackTool';
import { ListAllTracksTool } from './ListAllTracksTool';
export { AddNotesTool, RemoveNotesTool, ReadMusicTool, ReadChordProgressionTool, UpdateTodoListTool };
export {
AddNotesTool,
RemoveNotesTool,
ReadMusicTool,
ReadChordProgressionTool,
UpdateTodoListTool,
GetUserSelectedMusicRangeAndTrackTool,
ListAllTracksTool,
};
// Tool registry for easy access
export const AVAILABLE_TOOLS = {
@@ -19,6 +29,8 @@ export const AVAILABLE_TOOLS = {
remove_notes: RemoveNotesTool,
read_music: ReadMusicTool,
read_chord_progression: ReadChordProgressionTool,
get_user_selected_music_range_and_track: GetUserSelectedMusicRangeAndTrackTool,
list_all_tracks: ListAllTracksTool,
} as const;
export type ToolName = keyof typeof AVAILABLE_TOOLS;
+211
View File
@@ -0,0 +1,211 @@
import { KGCore } from '../../core/KGCore';
import { KGRegion } from '../../core/region/KGRegion';
import { KGMidiRegion } from '../../core/region/KGMidiRegion';
import { KGGlobalRegion } from '../../core/region/KGGlobalRegion';
import { KGMidiTrack } from '../../core/track/KGMidiTrack';
import { KGTrack } from '../../core/track/KGTrack';
import { useProjectStore } from '../../stores/projectStore';
export const NO_MIDI_TARGET_RAW_MESSAGE =
'No MIDI target could be resolved. Select the MIDI region you want me to edit and retry, or tell me which MIDI track to operate on by providing its track_id.';
export const NO_MIDI_TARGET_HISTORY_MESSAGE =
'I could not tell which MIDI content to edit. Select a MIDI region and retry, or tell me which track I should work on.';
export const NO_MIDI_TARGET_UI_MESSAGE =
'Select a MIDI region, or specify a track.';
export interface ActiveMidiRegionContext {
region: KGMidiRegion;
track: KGMidiTrack;
trackName: string;
}
export interface SelectedMusicRangeContext {
section: string;
startBeat: number | null;
endBeat: number | null;
hasRange: boolean;
}
export interface SelectedTrackContext {
track: KGTrack | null;
trackId: string | null;
trackName: string | null;
hasSelectedTrack: boolean;
}
export function getTrackDisplayName(track: { getName(): string; getTrackIndex(): number }): string {
return track.getName() || `Track ${track.getTrackIndex() + 1}`;
}
export function resolveMidiTrackByIdOrName(
trackId: string | undefined,
trackName: string | undefined,
): KGMidiTrack | null {
const project = KGCore.instance().getCurrentProject();
const midiTracks = project.getTracks().filter((track): track is KGMidiTrack => track instanceof KGMidiTrack);
if (trackId) {
return midiTracks.find(track => track.getId().toString() === trackId) ?? null;
}
if (trackName) {
return midiTracks.find(track => track.getName() === trackName) ?? null;
}
return null;
}
export function findRegionById(regionId: string): KGRegion | null {
const project = KGCore.instance().getCurrentProject();
for (const track of project.getTracks()) {
const region = track.getRegions().find(candidate => candidate.getId() === regionId);
if (region) {
return region;
}
}
for (const globalTrack of project.getGlobalTracks()) {
const region = globalTrack.getRegions().find(candidate => candidate.getId() === regionId);
if (region) {
return region;
}
}
return null;
}
export function findRegularTrackByRegion(region: KGRegion): KGTrack | null {
const project = KGCore.instance().getCurrentProject();
return project.getTracks().find(track => track.getRegions().includes(region)) ?? null;
}
export function resolveSelectedMusicRangeContext(): SelectedMusicRangeContext {
const project = KGCore.instance().getCurrentProject();
const storeState = useProjectStore.getState();
if (project.getIsLooping()) {
const beatsPerBar = project.getTimeSignature().numerator;
const [loopStartBar, loopEndBar] = project.getLoopingRange();
const startBeat = loopStartBar * beatsPerBar;
const endBeat = (loopEndBar + 1) * beatsPerBar;
return {
section: `- Start Beat: ${startBeat}\n- End Beat: ${endBeat}`,
startBeat,
endBeat,
hasRange: true,
};
}
const selectedRegions = (storeState.selectedRegionIds ?? [])
.map(regionId => findRegionById(regionId))
.filter((region): region is KGRegion => region !== null);
if (selectedRegions.length === 0) {
return {
section: '- No selected music range.',
startBeat: null,
endBeat: null,
hasRange: false,
};
}
const startBeat = Math.min(...selectedRegions.map(region => region.getStartFromBeat()));
const endBeat = Math.max(...selectedRegions.map(region => region.getStartFromBeat() + region.getLength()));
return {
section: `- Start Beat: ${startBeat}\n- End Beat: ${endBeat}`,
startBeat,
endBeat,
hasRange: true,
};
}
export function resolveSelectedTrackContext(): SelectedTrackContext {
const project = KGCore.instance().getCurrentProject();
const storeState = useProjectStore.getState();
const selectedRegionIds = storeState.selectedRegionIds ?? [];
if (selectedRegionIds.length === 0) {
return {
track: null,
trackId: null,
trackName: null,
hasSelectedTrack: false,
};
}
const selectedRegions = selectedRegionIds
.map(regionId => findRegionById(regionId))
.filter((region): region is KGRegion => region !== null);
if (selectedRegions.length === 0 || selectedRegions.every(region => region instanceof KGGlobalRegion)) {
return {
track: null,
trackId: null,
trackName: null,
hasSelectedTrack: false,
};
}
const selectedTrackId = storeState.selectedTrackId;
const selectedTrack = selectedTrackId
? project.getTracks().find(track => track.getId().toString() === selectedTrackId) ?? null
: null;
if (!selectedTrack) {
return {
track: null,
trackId: null,
trackName: null,
hasSelectedTrack: false,
};
}
return {
track: selectedTrack,
trackId: selectedTrack.getId().toString(),
trackName: getTrackDisplayName(selectedTrack),
hasSelectedTrack: true,
};
}
export function resolveActiveOrSelectedMidiRegionContext(): ActiveMidiRegionContext | null {
const project = KGCore.instance().getCurrentProject();
const tracks = project.getTracks();
const midiTracks = tracks.filter((track): track is KGMidiTrack => track instanceof KGMidiTrack);
const storeState = useProjectStore.getState();
if (storeState.activeRegionId) {
for (const track of midiTracks) {
const region = track.getRegions().find(candidate => candidate.getId() === storeState.activeRegionId);
if (region instanceof KGMidiRegion) {
return {
region,
track,
trackName: getTrackDisplayName(track),
};
}
}
}
const selectedItems = KGCore.instance().getSelectedItems();
for (const item of selectedItems) {
if (!(item instanceof KGMidiRegion)) {
continue;
}
const track = midiTracks.find(candidate => candidate.getId().toString() === item.getTrackId());
if (track) {
return {
region: item,
track,
trackName: getTrackDisplayName(track),
};
}
}
return null;
}
+10 -8
View File
@@ -341,13 +341,13 @@ export class KGDebugger {
* Usage examples in browser console:
*
* // Single tool call:
* await KGStudio.KGDebugger.testToolCall('{"name":"read_music","arguments":{"start":0,"length":8}}')
* await KGDebugger.testToolCall('{"name":"read_music","arguments":{"start":0,"length":8}}')
*
* // Multiple tool calls:
* await KGStudio.KGDebugger.testToolCall('[{"name":"remove_notes","arguments":{"start":0,"end_beat":4}},{"name":"add_notes","arguments":{"notes":[{"pitch":"C4","start":0,"length":1}]}}]')
* await KGDebugger.testToolCall('[{"name":"remove_notes","arguments":{"start":0,"end_beat":4}},{"name":"add_notes","arguments":{"notes":[{"pitch":"C4","start":0,"length":1}]}}]')
*
* // Can also pass a JS object directly (no need to stringify):
* await KGStudio.KGDebugger.testToolCall({name:"read_music",arguments:{start:0}})
* await KGDebugger.testToolCall({name:"read_music",arguments:{start:0}})
*
* @param input - JSON string, object, or array of tool call(s).
* Each tool call should have: { name: string, arguments: object }
@@ -430,13 +430,15 @@ export class KGDebugger {
console.log(" - Use browser developer tools for best experience");
console.log("");
console.log("💡 testToolCall examples:");
console.log(' await KGStudio.KGDebugger.testToolCall(\'{"name":"read_music","arguments":{"start":0,"length":8}}\')');
console.log(' await KGStudio.KGDebugger.testToolCall({name:"add_notes",arguments:{notes:[{pitch:"C4",start:0,length:1}]}})');
console.log(' await KGStudio.KGDebugger.testToolCall([{name:"remove_notes",arguments:{start:0,end_beat:4}},{name:"read_music",arguments:{}}])');
console.log(' await KGDebugger.testToolCall(\'{"name":"get_user_selected_music_range_and_track","arguments":{}}\')');
console.log(' await KGDebugger.testToolCall(\'{"name":"list_all_tracks","arguments":{}}\')');
console.log(' await KGDebugger.testToolCall(\'{"name":"read_music","arguments":{"start":0,"length":8}}\')');
console.log(' await KGDebugger.testToolCall({name:"add_notes",arguments:{notes:[{pitch:"C4",start:0,length:1}]}})');
console.log(' await KGDebugger.testToolCall([{name:"remove_notes",arguments:{start:0,end_beat:4}},{name:"read_music",arguments:{}}])');
console.log("");
console.log("💡 KGOne input examples:");
console.log(' await KGStudio.KGDebugger.inputKGOneCaption("Genre: Eurodance, 90s dance-pop, upbeat electronic...", 30)');
console.log(' await KGStudio.KGDebugger.inputKGOneLyrics("[Verse 1]\\nYour lyrics here...\\n\\n[Chorus]\\n...", 30)');
console.log(' await KGDebugger.inputKGOneCaption("Genre: Eurodance, 90s dance-pop, upbeat electronic...", 30)');
console.log(' await KGDebugger.inputKGOneLyrics("[Verse 1]\\nYour lyrics here...\\n\\n[Chorus]\\n...", 30)');
}
/**
+77 -22
View File
@@ -54,6 +54,7 @@ vi.mock('../utils/chatMessageUtils', () => ({
import { AgentCore } from '../agent/core/AgentCore';
import { KGCore } from '../core/KGCore';
import { KGMidiRegion } from '../core/region/KGMidiRegion';
import { KGMidiTrack } from '../core/track/KGMidiTrack';
import { useStreamProcessor } from './useStreamProcessor';
import type { ChatMessage } from '../types/projectTypes';
@@ -64,6 +65,7 @@ const flushMicrotasks = async (): Promise<void> => {
describe('useStreamProcessor', () => {
beforeEach(() => {
vi.restoreAllMocks();
mockedStoreState.activeRegionId = null;
mockedStoreState.toolFastForwardEnabled = false;
mockedStoreState.setToolFastForwardEnabled.mockClear();
@@ -297,16 +299,12 @@ describe('useStreamProcessor', () => {
},
} as unknown as AgentCore);
const leadTrack = new KGMidiTrack('Lead', 1);
leadTrack.setRegions([selectedRegion]);
vi.spyOn(KGCore, 'instance').mockReturnValue({
getCurrentProject: () => ({
getTimeSignature: () => ({ numerator: 4, denominator: 4 }),
getTracks: () => [
{
getId: () => '1',
getName: () => 'Lead',
getRegions: () => [selectedRegion],
},
],
getTracks: () => [leadTrack],
}),
getSelectedItems: () => [selectedRegion],
} as unknown as KGCore);
@@ -406,6 +404,71 @@ describe('useStreamProcessor', () => {
expect(toolResultMessage?.toolResultDisplayContent).toBe('raw music result');
});
it('uses tool-specific history and UI strings for error results when provided', async () => {
vi.spyOn(AgentCore, 'instance').mockReturnValue({
getAgentState: () => ({
getTodos: () => [],
}),
processUserInput: async function* () {
yield {
type: 'tool_call',
content: '',
toolCall: {
id: 'add-call-no-target',
type: 'function',
function: {
name: 'add_notes',
arguments: JSON.stringify({
notes: [{ pitch: 'C4', start: 0, length: 1 }],
}),
},
},
};
yield {
type: 'tool_result',
content: '',
toolResult: {
toolCallId: 'add-call-no-target',
name: 'add_notes',
success: false,
result: 'No MIDI target could be resolved. Select the MIDI region you want me to edit and retry, or tell me which MIDI track to operate on by providing its track_id.',
},
};
yield { type: 'done', content: '' };
},
} as unknown as AgentCore);
const messages = new Map<string, ChatMessage>();
const { result } = renderHook(() => useStreamProcessor({
onMessageAdd: (message) => {
messages.set(message.id, message);
},
onMessageUpdate: (messageId, updater) => {
const current = messages.get(messageId);
if (!current) {
throw new Error(`Missing message ${messageId}`);
}
messages.set(messageId, updater(current));
},
onMessageRemove: (messageId) => {
messages.delete(messageId);
},
onProcessingChange: () => undefined,
}));
await act(async () => {
await result.current.processStream('add notes prompt');
});
const addNotesMessage = [...messages.values()].find(message => message.toolName === 'add_notes');
expect(addNotesMessage?.toolRawResult).toBe(
'No MIDI target could be resolved. Select the MIDI region you want me to edit and retry, or tell me which MIDI track to operate on by providing its track_id.',
);
expect(addNotesMessage?.content).toContain('I could not tell which MIDI content to edit.');
expect(addNotesMessage?.toolResultDisplayContent).toBe('Select a MIDI region, or specify a track.');
});
it('shows a confirmation card and replaces it with a denied result when the user denies execution', async () => {
vi.spyOn(AgentCore, 'instance').mockReturnValue({
getAgentState: () => ({
@@ -440,16 +503,12 @@ describe('useStreamProcessor', () => {
} as unknown as AgentCore);
const selectedRegion = new KGMidiRegion('region-1', '1', 0, 'Verse Melody');
const leadTrack = new KGMidiTrack('Lead', 1);
leadTrack.setRegions([selectedRegion]);
vi.spyOn(KGCore, 'instance').mockReturnValue({
getCurrentProject: () => ({
getTimeSignature: () => ({ numerator: 4, denominator: 4 }),
getTracks: () => [
{
getId: () => '1',
getName: () => 'Lead',
getRegions: () => [selectedRegion],
},
],
getTracks: () => [leadTrack],
}),
getSelectedItems: () => [selectedRegion],
} as unknown as KGCore);
@@ -481,7 +540,7 @@ describe('useStreamProcessor', () => {
const confirmationMessage = [...messages.values()].find(message => message.toolConfirmation);
expect(confirmationMessage?.toolConfirmation?.toolName).toBe('add_notes');
expect(confirmationMessage?.toolConfirmation?.message).toContain('Allow creating 1 note in region **Verse Melody**');
expect(confirmationMessage?.toolConfirmation?.message).toContain('Allow creating 1 note on track **Lead** in region **Verse Melody**');
act(() => {
confirmationMessage?.onToolConfirmationDecision?.('deny');
@@ -533,16 +592,12 @@ describe('useStreamProcessor', () => {
} as unknown as AgentCore);
const selectedRegion = new KGMidiRegion('region-1', '1', 0, 'Intro');
const leadTrack = new KGMidiTrack('Lead', 1);
leadTrack.setRegions([selectedRegion]);
vi.spyOn(KGCore, 'instance').mockReturnValue({
getCurrentProject: () => ({
getTimeSignature: () => ({ numerator: 4, denominator: 4 }),
getTracks: () => [
{
getId: () => '1',
getName: () => 'Lead',
getRegions: () => [selectedRegion],
},
],
getTracks: () => [leadTrack],
}),
getSelectedItems: () => [selectedRegion],
} as unknown as KGCore);
+15 -11
View File
@@ -174,17 +174,21 @@ export const useStreamProcessor = (options: StreamProcessorOptions): StreamProce
const pendingToolCall = pendingToolCallIndex >= 0
? pendingToolCalls.splice(pendingToolCallIndex, 1)[0]
: undefined;
let toolHistoryContent = result;
let toolResultDisplayContent = result;
if (success) {
try {
const toolInstance = createToolInstance(name);
toolResultDisplayContent = toolInstance?.buildToolResultDisplayContent(
pendingToolCall?.arguments ?? null,
{ success, result },
) ?? result;
} catch {
toolResultDisplayContent = result;
}
try {
const toolInstance = createToolInstance(name);
toolHistoryContent = toolInstance?.buildToolHistoryContent(
pendingToolCall?.arguments ?? null,
{ success, result },
) ?? result;
toolResultDisplayContent = toolInstance?.buildToolResultDisplayContent(
pendingToolCall?.arguments ?? null,
{ success, result },
) ?? result;
} catch {
toolHistoryContent = result;
toolResultDisplayContent = result;
}
const toolResultMsg = name === TODO_TOOL_NAME
? {
@@ -194,7 +198,7 @@ export const useStreamProcessor = (options: StreamProcessorOptions): StreamProce
todoSnapshot: AgentCore.instance().getAgentState().getTodos().map(todo => ({ ...todo })),
}
: {
...createMessage('assistant', `${success ? '✅' : '❌'} **${name}**\n\n └── ${result}`),
...createMessage('assistant', `${success ? '✅' : '❌'} **${name}**\n\n └── ${toolHistoryContent}`),
toolName: name,
toolSuccess: success,
toolRawResult: result,
@@ -68,8 +68,10 @@ describe('abcNotationUtil - Integration Tests with Real Project Data', () => {
const result = convertRegionToABCNotation(melodyRegion, 0, 32);
// Expected ABC notation output (exactly as provided)
const expectedABCNotation = `X:1
T:Melody Region 1
const expectedABCNotation = `track_id: 1
track_name: Melody
Instrument: Acoustic Grand Piano
X:1
M:4/4
L:1/4
Q:1/4=125
@@ -95,10 +97,13 @@ E E F G | G F E D | C C D E | E3/2 D1/2 D2 | E E F G | G F E D | C C D E | D3/2
// Verify it contains rest notation and proper headers
expect(result).toBeDefined();
expect(result).toContain('T:Pad Chord Region 1');
expect(result).toContain('track_id: 2');
expect(result).toContain('track_name: Pad Chord');
expect(result).toContain('Instrument: Pad 1 (new age)');
expect(result).toContain('M:4/4');
expect(result).toContain('Q:1/4=125');
expect(result).toContain('K:C');
expect(result).not.toContain('\nT:');
expect(result).toContain('z'); // Should contain rest notation
});
@@ -167,8 +172,12 @@ E E F G | G F E D | C C D E | E3/2 D1/2 D2 | E E F G | G F E D | C C D E | D3/2
const melodyABC = convertRegionToABCNotation(melodyRegion, 0, 32);
const padABC = convertRegionToABCNotation(padRegion, 0, 32);
expect(melodyABC).toContain('T:Melody Region 1');
expect(padABC).toContain('T:Pad Chord Region 1');
expect(melodyABC).toContain('track_name: Melody');
expect(padABC).toContain('track_name: Pad Chord');
expect(melodyABC).toContain('Instrument: Acoustic Grand Piano');
expect(padABC).toContain('Instrument: Pad 1 (new age)');
expect(melodyABC).not.toContain('\nT:');
expect(padABC).not.toContain('\nT:');
});
it('should verify complete deserialization hierarchy', () => {
@@ -211,7 +220,8 @@ E E F G | G F E D | C C D E | E3/2 D1/2 D2 | E E F G | G F E D | C C D E | D3/2
const result = convertRegionToABCNotation(melodyRegion, 0, 8);
expect(result).toBeDefined();
expect(result).toContain('T:Melody Region 1');
expect(result).toContain('track_name: Melody');
expect(result).not.toContain('\nT:');
// Should only contain the first 2 bars of music
expect(result).not.toContain('D3/2 C1/2 C2'); // This appears later in the song
});
@@ -224,7 +234,8 @@ E E F G | G F E D | C C D E | E3/2 D1/2 D2 | E E F G | G F E D | C C D E | D3/2
const result = convertRegionToABCNotation(melodyRegion, 16, 24);
expect(result).toBeDefined();
expect(result).toContain('T:Melody Region 1');
expect(result).toContain('track_name: Melody');
expect(result).not.toContain('\nT:');
// Should start from the second repetition
const lines = result.split('\n');
const musicLine = lines[lines.length - 1]; // Last line contains the music
+26 -5
View File
@@ -8,6 +8,7 @@ import { KGMidiNote } from '../core/midi/KGMidiNote';
import { KGCore } from '../core/KGCore';
import { KGProject } from '../core/KGProject';
import { KGChordRegion } from '../core/region/KGChordRegion';
import { FLUIDR3_INSTRUMENT_MAP } from '../constants/generalMidiConstants';
import { pitchToNoteName } from './midiUtil';
import { beatsToTicks, getTicksPerBar, reduceFraction } from './mathUtil';
import type { TimeSignature } from '../types/projectTypes';
@@ -168,18 +169,38 @@ function convertTicksToABCLength(ticks: number, timeSignature: TimeSignature): s
* @param project - Project containing tempo and time signature info
* @returns ABC header string
*/
function resolveRegionTrackMetadata(region: KGMidiRegion, project: KGProject): {
trackId: string;
trackName: string;
instrumentName: string;
} {
const trackId = region.getTrackId();
const track = project.getTracks().find(candidate => candidate.getId().toString() === trackId);
const trackName = track?.getName() || 'Unnamed Track';
const instrumentKey = 'getInstrument' in (track ?? {}) && typeof track.getInstrument === 'function'
? track.getInstrument()
: null;
const instrumentName = instrumentKey
? FLUIDR3_INSTRUMENT_MAP[instrumentKey]?.displayName || instrumentKey
: 'Unknown Instrument';
return { trackId, trackName, instrumentName };
}
function formatABCHeader(region: KGMidiRegion, project: KGProject): string {
const timeSignature = project.getTimeSignature();
const bpm = project.getBpm();
const keySignature = project.getKeySignature();
const regionName = region.getName();
const { trackId, trackName, instrumentName } = resolveRegionTrackMetadata(region, project);
// Get ABC notation key signature from the key signature map
const abcKeySignature = KEY_SIGNATURE_MAP[keySignature]?.abcNotationKeySignature || 'C';
const header = [
`track_id: ${trackId}`,
`track_name: ${trackName}`,
`Instrument: ${instrumentName}`,
'X:1', // Reference number
`T:${regionName}`, // Title
`M:${timeSignature.numerator}/${timeSignature.denominator}`, // Time signature
`L:1/${timeSignature.denominator}`, // note length unit should be aligned with time signature
`Q:1/${timeSignature.denominator}=${bpm}`, // Tempo (quarter note = BPM)
@@ -353,12 +374,12 @@ export function convertBeatRangeChordProgressionToABCNotation(
if (segments.length === 0) {
return [
'Selected Region Chord Progression',
'Chord Progression',
'This progression comes only from user-defined chord regions on the global chord track. If no chord progression is defined for this range, read the notes directly with `read_music`.',
'',
header,
'',
'No chord progression is defined for the selected MIDI region range. Use `read_music` to inspect the notes directly.'
'No chord progression is defined for the requested range on the global chord track. Use `read_music` to inspect the notes directly.'
].join('\n');
}
@@ -367,7 +388,7 @@ export function convertBeatRangeChordProgressionToABCNotation(
const chordNotes = formatChordProgressionNoteLine(segments, timeSignature);
return [
'Selected Region Chord Progression',
'Chord Progression',
'This progression comes only from user-defined chord regions on the global chord track. Representation 1 uses symbolic chord names such as `Em7b5`. Representation 2 rewrites the same progression as note-based ABC chord tokens.',
'',
header,
@@ -35,12 +35,6 @@ vi.mock('../../stores/projectStore', () => ({
},
}));
vi.mock('../../agent/core/SystemPrompts', () => ({
SystemPrompts: {
getPromptWithContext: vi.fn(async (value: string) => value),
},
}));
vi.mock('../localLLMConfig', async () => {
const actual = await vi.importActual<typeof import('../localLLMConfig')>('../localLLMConfig');
return {
@@ -336,6 +330,17 @@ describe('processUserMessage slash commands', () => {
expect(result.metadata).toMatchObject({ error: 'local_browser_unsupported' });
});
it('passes non-command messages through to the LLM even when no region is selected', async () => {
const result = await processUserMessage('hello');
expect(result.sendToLLM).toBe(true);
expect(result.finalMessageForLLM).toBe('hello');
expect(result.pseudoAssistantResponse).toBeNull();
expect(result.metadata).toMatchObject({
mode: 'pass_through_plain_user_message',
});
});
it('allows local-browser messages when only SharedArrayBuffer isolation support is missing', async () => {
detectLocalLLMRuntimeSupportMock.mockReturnValue({
supported: true,
@@ -345,12 +350,11 @@ describe('processUserMessage slash commands', () => {
secureContext: true,
reason: 'This host may not support the local browser runtime reliably because cross-origin isolation or SharedArrayBuffer is unavailable. COOP/COEP headers may be missing.',
});
storeState.activeRegionId = 'region-1';
const result = await processUserMessage('hello');
expect(result.sendToLLM).toBe(true);
expect(result.finalMessageForLLM).toContain('hello');
expect(result.finalMessageForLLM).toBe('hello');
expect(result.pseudoAssistantResponse).toBeNull();
});
});
+3 -45
View File
@@ -1,7 +1,6 @@
import { clearChatHistoryAndUI } from '../chatUtil';
import { useProjectStore } from '../../stores/projectStore';
import { ConfigManager } from '../../core/config/ConfigManager';
import { SystemPrompts } from '../../agent/core/SystemPrompts';
import { detectLocalLLMRuntimeSupport, LOCAL_LLM_PROVIDER_KEY } from '../localLLMConfig';
import { normalizeLanguageSetting, resolveLanguageSetting } from '../../i18n/locale';
import type { ResolvedLocaleCode } from '../../i18n/types';
@@ -246,7 +245,7 @@ export async function processUserMessage(originalMessage: string): Promise<UserM
}
}
// Non-command message: require an active or selected region
// Non-command message: pass through to the LLM unchanged.
try {
// Provider-specific configuration checks before sending to LLM
try {
@@ -329,53 +328,12 @@ export async function processUserMessage(originalMessage: string): Promise<UserM
console.warn('Provider config check failed; proceeding with defaults', e);
}
const { activeRegionId, selectedRegionIds } = useProjectStore.getState();
const hasContextRegion = !!activeRegionId || (Array.isArray(selectedRegionIds) && selectedRegionIds.length > 0);
if (!hasContextRegion) {
// No region context: show guidance and do not send to LLM
const url = `${import.meta.env.BASE_URL}chat/error_no_selected_region.md`;
let md = 'Please select a region or open a MIDI region in the piano roll before asking for editing.';
try {
const resp = await fetch(url);
if (resp.ok) {
md = await resp.text();
}
} catch (e) {
console.warn('Failed to fetch error_no_selected_region.md', e);
// ignore fetch failure, use fallback text
}
return {
displayUserMessage: true,
sendToLLM: false,
finalMessageForLLM: null,
pseudoAssistantResponse: md,
metadata: { error: 'no_selected_region' }
};
}
// Has region context: pass through, but append processed appendix to the LLM-bound message
let appendix = '';
try {
const resp = await fetch(`${import.meta.env.BASE_URL}prompts/user_msg_appendix.md`);
if (resp.ok) {
const rawAppendix = await resp.text();
appendix = await SystemPrompts.getPromptWithContext(rawAppendix);
}
} catch (e) {
// If appendix fetch fails, proceed without it
console.warn('Failed to fetch user_msg_appendix.md', e);
}
const finalForLLM = appendix ? `${trimmed}${appendix}` : trimmed;
return {
displayUserMessage: true,
sendToLLM: true,
finalMessageForLLM: finalForLLM,
finalMessageForLLM: trimmed,
pseudoAssistantResponse: null,
metadata: { mode: 'pass_through_with_region', appendixIncluded: appendix.length > 0 }
metadata: { mode: 'pass_through_plain_user_message' }
};
} catch {
// Fallback: if store access fails, pass through unchanged