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
+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;
}