From a76a8abb940dbc00eb03b008241e14b69747c3ed Mon Sep 17 00:00:00 2001 From: Xiaohan-Tian <157918347+Xiaohan-Tian@users.noreply.github.com> Date: Sat, 6 Jun 2026 17:27:04 -0700 Subject: [PATCH] fix: support numeric track_id arguments across track tools --- package-lock.json | 4 +- src/agent/tools/AddNotesTool.test.ts | 63 ++++++++++++++++++++++ src/agent/tools/AddNotesTool.ts | 13 +++-- src/agent/tools/DeleteTrackTool.test.ts | 35 ++++++++++++ src/agent/tools/DeleteTrackTool.ts | 18 ++++--- src/agent/tools/ReadMusicTool.test.ts | 41 ++++++++++++++ src/agent/tools/ReadMusicTool.ts | 19 ++++--- src/agent/tools/RemoveNotesTool.test.ts | 72 +++++++++++++++++++++++++ src/agent/tools/RemoveNotesTool.ts | 15 +++--- src/agent/tools/UpdateTrackTool.test.ts | 42 +++++++++++++++ src/agent/tools/UpdateTrackTool.ts | 21 ++++---- src/agent/tools/trackIdNormalization.ts | 16 ++++++ 12 files changed, 322 insertions(+), 37 deletions(-) create mode 100644 src/agent/tools/trackIdNormalization.ts diff --git a/package-lock.json b/package-lock.json index bc5838f..50cb075 100644 --- a/package-lock.json +++ b/package-lock.json @@ -1,12 +1,12 @@ { "name": "K.G.Studio", - "version": "0.19.0-build.20260531", + "version": "0.20.2-build.20260606", "lockfileVersion": 3, "requires": true, "packages": { "": { "name": "K.G.Studio", - "version": "0.19.0-build.20260531", + "version": "0.20.2-build.20260606", "dependencies": { "@breezystack/lamejs": "^1.2.7", "class-transformer": "^0.5.1", diff --git a/src/agent/tools/AddNotesTool.test.ts b/src/agent/tools/AddNotesTool.test.ts index e1b23d0..815aed8 100644 --- a/src/agent/tools/AddNotesTool.test.ts +++ b/src/agent/tools/AddNotesTool.test.ts @@ -87,6 +87,26 @@ describe('AddNotesTool', () => { expect(createdRegion.getNotes().map(note => note.getStartBeat())).toEqual([0, 2]); }); + it('accepts a numeric track_id and creates notes on the matching track', async () => { + const track = new KGMidiTrack('Lead', 1); + const project = new KGProject('numeric-track-id-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(), + notes: [{ pitch: 'C4', start: 8, length: 2 }], + }); + + expect(result.success).toBe(true); + expect(track.getRegions()).toHaveLength(1); + const createdRegion = track.getRegions()[0] as KGMidiRegion; + expect(createdRegion.getName()).toBe('Lead Region'); + expect(createdRegion.getNotes()).toHaveLength(1); + expect(createdRegion.getNotes()[0].getStartBeat()).toBe(0); + }); + 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); @@ -126,6 +146,26 @@ describe('AddNotesTool', () => { expect((bassTrack.getRegions()[0] as KGMidiRegion).getName()).toBe('Bass Region'); }); + it('uses a numeric 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('numeric-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(), + 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); @@ -169,6 +209,29 @@ describe('AddNotesTool', () => { expect(regionB.getNotes().find(note => note.getId() !== 'note-existing')?.getStartBeat()).toBe(0); }); + it('builds summaries when track_id is provided as a number', () => { + const track = new KGMidiTrack('Lead', 1); + const project = new KGProject('numeric-track-id-summary-project', 8, 0, 120, { numerator: 4, denominator: 4 }, 'C major'); + project.setTracks([track]); + mockCore(project); + + const tool = new AddNotesTool(); + const args = { + track_id: track.getId(), + notes: [ + { pitch: 'C4', start: 8, length: 2 }, + { pitch: 'E4', start: 10, length: 2 }, + ], + }; + + expect(tool.buildToolResultDisplayContent(args, { success: true, result: 'raw result' })).toBe( + 'Successfully created 2 notes in new region **Lead Region** on track **Lead**, spanning bars 3 to 3.', + ); + expect(tool.buildConfirmationContent(args)).toBe( + 'Allow creating 2 notes on track **Lead** in a new region, spanning bars 3 to 3?', + ); + }); + 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); diff --git a/src/agent/tools/AddNotesTool.ts b/src/agent/tools/AddNotesTool.ts index ef273ed..1e609a8 100644 --- a/src/agent/tools/AddNotesTool.ts +++ b/src/agent/tools/AddNotesTool.ts @@ -16,6 +16,7 @@ import { ResizeRegionCommand } from '../../core/commands/region/ResizeRegionComm import { KGCore } from '../../core/KGCore'; import { KGMidiRegion } from '../../core/region/KGMidiRegion'; import { KGMidiTrack } from '../../core/track/KGMidiTrack'; +import { normalizeOptionalTrackIdParam } from './trackIdNormalization'; interface RequestedNote { pitch: string; @@ -226,9 +227,10 @@ export class AddNotesTool extends BaseTool { async execute(params: Record): Promise { try { - this.validateParameters(params); + const normalizedParams = normalizeOptionalTrackIdParam(params); + this.validateParameters(normalizedParams); - const notes = params.notes as RequestedNote[]; + const notes = normalizedParams.notes as RequestedNote[]; if (notes.length === 0) { return this.createErrorResult('No notes were provided.'); } @@ -255,8 +257,8 @@ export class AddNotesTool extends BaseTool { } } - const trackId = params.track_id as string | undefined; - const trackName = params.track_name as string | undefined; + const trackId = normalizedParams.track_id as string | undefined; + const trackName = normalizedParams.track_name as string | undefined; if (trackId || trackName) { const explicitTrack = resolveMidiTrackByIdOrName(trackId, trackName); if (!explicitTrack) { @@ -291,7 +293,8 @@ export class AddNotesTool extends BaseTool { } private buildSummaryData(args: Record): AddNotesSummaryData | null { - const typedArgs = args as { + const normalizedArgs = normalizeOptionalTrackIdParam(args); + const typedArgs = normalizedArgs as { notes?: Array<{ start: number; length: number }>; track_id?: string; track_name?: string; diff --git a/src/agent/tools/DeleteTrackTool.test.ts b/src/agent/tools/DeleteTrackTool.test.ts index 046655f..c56cb38 100644 --- a/src/agent/tools/DeleteTrackTool.test.ts +++ b/src/agent/tools/DeleteTrackTool.test.ts @@ -52,6 +52,24 @@ describe('DeleteTrackTool', () => { ); }); + it('deletes a MIDI track by numeric track_id', async () => { + const leadTrack = new KGMidiTrack('Lead', 1, 'trumpet'); + const bassTrack = new KGMidiTrack('Bass', 2, 'acoustic_bass'); + const project = new KGProject('delete-by-numeric-id-project'); + project.setTracks([leadTrack, bassTrack]); + mockCore(project); + + const tool = new DeleteTrackTool(); + const result = await tool.execute({ track_id: 1 }); + + expect(result).toEqual({ + success: true, + result: 'Track deleted:\ntrack_id: 1\ntrack_name: Lead', + }); + expect(project.getTracks().map(track => track.getName())).toEqual(['Bass']); + expect(tool.buildConfirmationContent({ track_id: 1 })).toBe('Allow deleting track ID **1**?'); + }); + it('deletes a MIDI track by track_name', async () => { const leadTrack = new KGMidiTrack('Lead', 1, 'trumpet'); const bassTrack = new KGMidiTrack('Bass', 2, 'acoustic_bass'); @@ -86,6 +104,23 @@ describe('DeleteTrackTool', () => { expect(project.getTracks().map(track => track.getName())).toEqual(['Lead']); }); + it('uses numeric track_id when both track_id and track_name are provided', async () => { + const leadTrack = new KGMidiTrack('Lead', 1, 'trumpet'); + const bassTrack = new KGMidiTrack('Bass', 2, 'acoustic_bass'); + const project = new KGProject('delete-numeric-track-id-precedence-project'); + project.setTracks([leadTrack, bassTrack]); + mockCore(project); + + const tool = new DeleteTrackTool(); + const result = await tool.execute({ + track_id: 2, + track_name: 'Lead', + }); + + expect(result.success).toBe(true); + expect(project.getTracks().map(track => track.getName())).toEqual(['Lead']); + }); + it('rejects duplicate track names when track_id is omitted', async () => { const firstLead = new KGMidiTrack('Lead', 1, 'trumpet'); const secondLead = new KGMidiTrack('Lead', 2, 'flute'); diff --git a/src/agent/tools/DeleteTrackTool.ts b/src/agent/tools/DeleteTrackTool.ts index f4bfe8c..e77236b 100644 --- a/src/agent/tools/DeleteTrackTool.ts +++ b/src/agent/tools/DeleteTrackTool.ts @@ -3,6 +3,7 @@ import { KGMidiTrack } from '../../core/track/KGMidiTrack'; import { BaseTool } from './BaseTool'; import type { ToolParameter, ToolResult } from './BaseTool'; import { resolveMidiTrackByExactName, resolveMidiTrackByIdOrName } from './toolTargeting'; +import { normalizeOptionalTrackIdParam } from './trackIdNormalization'; export class DeleteTrackTool extends BaseTool { readonly name = 'delete_track'; @@ -35,12 +36,14 @@ export class DeleteTrackTool extends BaseTool { return undefined; } - if (typeof args.track_id === 'string') { - return `Allow deleting track ID **${args.track_id}**?`; + const normalizedArgs = normalizeOptionalTrackIdParam(args); + + if (typeof normalizedArgs.track_id === 'string') { + return `Allow deleting track ID **${normalizedArgs.track_id}**?`; } - if (typeof args.track_name === 'string') { - return `Allow deleting track **${args.track_name}**?`; + if (typeof normalizedArgs.track_name === 'string') { + return `Allow deleting track **${normalizedArgs.track_name}**?`; } return undefined; @@ -66,10 +69,11 @@ export class DeleteTrackTool extends BaseTool { async execute(params: Record): Promise { try { - this.validateParameters(params); + const normalizedParams = normalizeOptionalTrackIdParam(params); + this.validateParameters(normalizedParams); - const trackId = params.track_id as string | undefined; - const trackName = params.track_name as string | undefined; + const trackId = normalizedParams.track_id as string | undefined; + const trackName = normalizedParams.track_name as string | undefined; if (!trackId && !trackName) { return this.createErrorResult('Either track_id or track_name must be provided.'); diff --git a/src/agent/tools/ReadMusicTool.test.ts b/src/agent/tools/ReadMusicTool.test.ts index cd87183..7aaa5aa 100644 --- a/src/agent/tools/ReadMusicTool.test.ts +++ b/src/agent/tools/ReadMusicTool.test.ts @@ -68,6 +68,29 @@ describe('ReadMusicTool', () => { expect(result.result).toContain('z4 | z4 | // No regions found'); }); + it('reads a specific track when track_id is provided as a number', async () => { + const project = new KGProject('read-music-numeric-track-project', 8, 0, 120, { numerator: 4, denominator: 4 }, 'C major'); + const leadTrack = buildTrack('Lead', 1, 0, 8); + const bassTrack = buildTrack('Bass', 2, 0, 8); + project.setTracks([leadTrack, bassTrack]); + + vi.spyOn(KGCore, 'instance').mockReturnValue({ + getCurrentProject: () => project, + } as unknown as KGCore); + + const tool = new ReadMusicTool(); + const result = await tool.execute({ track_id: 2, start: 0, length: 8 }); + + expect(result.success).toBe(true); + expect(result.result).toContain('track_id: 2'); + expect(result.result).toContain('track_name: Bass'); + expect(result.result).not.toContain('track_id: 1'); + expect(tool.buildToolResultDisplayContent( + { track_id: 2, start: 0, length: 8 }, + { success: true, result: 'raw result' }, + )).toBe('Read track Bass from bars 1 to 2.'); + }); + it('returns a professional empty-project message when all MIDI tracks are empty', async () => { const project = new KGProject('read-music-project', 8, 0, 120, { numerator: 4, denominator: 4 }, 'C major'); const emptyTrack = new KGMidiTrack('Lead', 1); @@ -87,6 +110,24 @@ describe('ReadMusicTool', () => { expect(result.result).toBe('No musical content is present in the project yet.'); }); + it('preserves the all-tracks behavior when track_id is "all"', async () => { + const project = new KGProject('read-music-all-track-project', 8, 0, 120, { numerator: 4, denominator: 4 }, 'C major'); + const leadTrack = buildTrack('Lead', 1, 0, 8); + const bassTrack = buildTrack('Bass', 2, 0, 8); + project.setTracks([leadTrack, bassTrack]); + + vi.spyOn(KGCore, 'instance').mockReturnValue({ + getCurrentProject: () => project, + } as unknown as KGCore); + + const tool = new ReadMusicTool(); + const result = await tool.execute({ track_id: 'all', start: 0, length: 8 }); + + expect(result.success).toBe(true); + expect(result.result).toContain('track_id: 1'); + expect(result.result).toContain('track_id: 2'); + }); + 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); diff --git a/src/agent/tools/ReadMusicTool.ts b/src/agent/tools/ReadMusicTool.ts index 647a87b..e66dea3 100644 --- a/src/agent/tools/ReadMusicTool.ts +++ b/src/agent/tools/ReadMusicTool.ts @@ -6,6 +6,7 @@ import { KGMidiRegion } from '../../core/region/KGMidiRegion'; import { convertRegionToABCNotation } from '../../util/abcNotationUtil'; import { KEY_SIGNATURE_MAP } from '../../constants/coreConstants'; import { FLUIDR3_INSTRUMENT_MAP } from '../../constants/generalMidiConstants'; +import { normalizeOptionalTrackIdParam } from './trackIdNormalization'; /** * Tool for reading music content from the project @@ -68,12 +69,13 @@ export class ReadMusicTool extends BaseTool { async execute(params: Record): Promise { try { + const normalizedParams = normalizeOptionalTrackIdParam(params); // Validate parameters - this.validateParameters(params); + this.validateParameters(normalizedParams); - const trackId = params.track_id as string | undefined; - const startBeat = (params.start as number) || 0; - const length = params.length as number | undefined; + const trackId = normalizedParams.track_id as string | undefined; + const startBeat = (normalizedParams.start as number) || 0; + const length = normalizedParams.length as number | undefined; const project = this.getCurrentProject(); const tracks = project.getTracks(); @@ -140,6 +142,7 @@ export class ReadMusicTool extends BaseTool { startBar: number; endBar: number; } | null { + const normalizedArgs = normalizeOptionalTrackIdParam(args); const project = this.getCurrentProject(); const tracks = project.getTracks(); if (tracks.length === 0) { @@ -147,8 +150,8 @@ export class ReadMusicTool extends BaseTool { } const beatsPerBar = project.getTimeSignature().numerator; - const startBeat = (args.start as number) || 0; - const length = args.length as number | undefined; + const startBeat = (normalizedArgs.start as number) || 0; + const length = normalizedArgs.length as number | undefined; if (startBeat < 0 || (length !== undefined && length <= 0)) { return null; } @@ -157,9 +160,9 @@ export class ReadMusicTool extends BaseTool { const rawEndBeat = length !== undefined ? startBeat + length : undefined; const roundedEndBeat = rawEndBeat !== undefined ? Math.ceil(rawEndBeat / beatsPerBar) * beatsPerBar - : this.getTrackReadEndBeat(args, tracks, roundedStartBeat); + : this.getTrackReadEndBeat(normalizedArgs, tracks, roundedStartBeat); - const trackNames = this.resolveSummaryTrackNames(args, tracks); + const trackNames = this.resolveSummaryTrackNames(normalizedArgs, tracks); if (trackNames.length === 0 || roundedEndBeat === undefined) { return null; } diff --git a/src/agent/tools/RemoveNotesTool.test.ts b/src/agent/tools/RemoveNotesTool.test.ts index 8b23d07..74a42fc 100644 --- a/src/agent/tools/RemoveNotesTool.test.ts +++ b/src/agent/tools/RemoveNotesTool.test.ts @@ -82,6 +82,29 @@ describe('RemoveNotesTool', () => { expect(regionB.getNotes()).toHaveLength(0); }); + it('removes notes across all MIDI regions on a track when track_id is numeric', 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('numeric-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(), + 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); @@ -133,6 +156,32 @@ describe('RemoveNotesTool', () => { expect(bassRegion.getNotes()).toHaveLength(1); }); + it('uses numeric 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('numeric-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(), + 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); @@ -158,6 +207,29 @@ describe('RemoveNotesTool', () => { expect(secondRegion.getNotes()).toHaveLength(1); }); + it('builds confirmation and result summaries when track_id is numeric', () => { + const track = new KGMidiTrack('Lead', 1); + const region = new KGMidiRegion('region-1', track.getId().toString(), track.getTrackIndex(), 'Verse Melody', 0, 32); + region.setNotes([ + new KGMidiNote('note-1', 16, 20, 60, 100), + new KGMidiNote('note-2', 20, 28, 64, 100), + ]); + track.setRegions([region]); + const project = new KGProject('numeric-remove-summary-project', 8, 0, 120, { numerator: 4, denominator: 4 }, 'C major'); + project.setTracks([track]); + mockCore(project); + + const tool = new RemoveNotesTool(); + const args = { track_id: track.getId(), start: 16, end: 24 }; + + expect(tool.buildConfirmationContent(args)).toBe( + 'Allow removing 2 notes from beats 16-24, 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, on track **Lead**, spanning bars 5 to 7.', + ); + }); + 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); diff --git a/src/agent/tools/RemoveNotesTool.ts b/src/agent/tools/RemoveNotesTool.ts index 5cfa4f9..0972bae 100644 --- a/src/agent/tools/RemoveNotesTool.ts +++ b/src/agent/tools/RemoveNotesTool.ts @@ -12,6 +12,7 @@ import { DeleteNotesCommand } from '../../core/commands/note/DeleteNotesCommand' import { KGMidiNote } from '../../core/midi/KGMidiNote'; import { KGMidiRegion } from '../../core/region/KGMidiRegion'; import { KGMidiTrack } from '../../core/track/KGMidiTrack'; +import { normalizeOptionalTrackIdParam } from './trackIdNormalization'; interface RemoveTargetRegionContext { region: KGMidiRegion; @@ -106,12 +107,13 @@ export class RemoveNotesTool extends BaseTool { async execute(params: Record): Promise { try { - this.validateParameters(params); + const normalizedParams = normalizeOptionalTrackIdParam(params); + this.validateParameters(normalizedParams); - const startBeat = params.start as number; - const endBeat = params.end as number; - const trackId = params.track_id as string | undefined; - const trackName = params.track_name as string | undefined; + const startBeat = normalizedParams.start as number; + const endBeat = normalizedParams.end as number; + const trackId = normalizedParams.track_id as string | undefined; + const trackName = normalizedParams.track_name as string | undefined; if (startBeat < 0) { return this.createErrorResult(`Invalid start ${startBeat}. Must be >= 0.`); @@ -159,7 +161,8 @@ export class RemoveNotesTool extends BaseTool { } private buildSummaryData(args: Record): RemoveNotesSummaryData | null { - const typedArgs = args as { + const normalizedArgs = normalizeOptionalTrackIdParam(args); + const typedArgs = normalizedArgs as { start?: number; end?: number; track_id?: string; diff --git a/src/agent/tools/UpdateTrackTool.test.ts b/src/agent/tools/UpdateTrackTool.test.ts index 8fd099b..aa0c5d5 100644 --- a/src/agent/tools/UpdateTrackTool.test.ts +++ b/src/agent/tools/UpdateTrackTool.test.ts @@ -56,6 +56,29 @@ describe('UpdateTrackTool', () => { ); }); + it('renames a track by numeric track_id', async () => { + const track = new KGMidiTrack('Lead', 1, 'trumpet'); + const project = new KGProject('rename-track-numeric-project'); + project.setTracks([track]); + mockCore(project); + + const tool = new UpdateTrackTool(); + const result = await tool.execute({ + track_id: 1, + new_track_name: 'Lead 2', + }); + + expect(result).toEqual({ + success: true, + result: 'Track updated:\ntrack_id: 1\ntrack_name: Lead 2\ninstrument: Trumpet', + }); + expect(track.getName()).toBe('Lead 2'); + expect(tool.buildConfirmationContent({ + track_id: 1, + new_track_name: 'Lead 2', + })).toBe('Allow updating track ID **1** to rename to **Lead 2**?'); + }); + it('updates a track instrument by track_name', async () => { const track = new KGMidiTrack('Lead', 1, 'trumpet'); const project = new KGProject('instrument-track-project'); @@ -115,6 +138,25 @@ describe('UpdateTrackTool', () => { expect(bassTrack.getName()).toBe('Bass 2'); }); + it('uses numeric track_id when both track_id and track_name are provided', async () => { + const leadTrack = new KGMidiTrack('Lead', 1, 'trumpet'); + const bassTrack = new KGMidiTrack('Bass', 2, 'acoustic_bass'); + const project = new KGProject('numeric-track-id-precedence-project'); + project.setTracks([leadTrack, bassTrack]); + mockCore(project); + + const tool = new UpdateTrackTool(); + const result = await tool.execute({ + track_id: 2, + track_name: 'Lead', + new_track_name: 'Bass 2', + }); + + expect(result.success).toBe(true); + expect(leadTrack.getName()).toBe('Lead'); + expect(bassTrack.getName()).toBe('Bass 2'); + }); + it('rejects duplicate track names when track_id is omitted', async () => { const firstLead = new KGMidiTrack('Lead', 1, 'trumpet'); const secondLead = new KGMidiTrack('Lead', 2, 'flute'); diff --git a/src/agent/tools/UpdateTrackTool.ts b/src/agent/tools/UpdateTrackTool.ts index cea35e4..43bf5ed 100644 --- a/src/agent/tools/UpdateTrackTool.ts +++ b/src/agent/tools/UpdateTrackTool.ts @@ -9,6 +9,7 @@ import { resolveMidiTrackByExactName, resolveMidiTrackByIdOrName, } from './toolTargeting'; +import { normalizeOptionalTrackIdParam } from './trackIdNormalization'; export class UpdateTrackTool extends BaseTool { readonly name = 'update_track'; @@ -51,12 +52,13 @@ export class UpdateTrackTool extends BaseTool { return undefined; } + const normalizedArgs = normalizeOptionalTrackIdParam(args); const normalizedInstrumentName = this.normalizeOptionalString(args.instrument); const normalizedNewTrackName = this.normalizeOptionalString(args.new_track_name); - const targetLabel = typeof args.track_id === 'string' - ? `track ID **${args.track_id}**` - : typeof args.track_name === 'string' - ? `track **${args.track_name}**` + const targetLabel = typeof normalizedArgs.track_id === 'string' + ? `track ID **${normalizedArgs.track_id}**` + : typeof normalizedArgs.track_name === 'string' + ? `track **${normalizedArgs.track_name}**` : null; if (!targetLabel) { return undefined; @@ -98,12 +100,13 @@ export class UpdateTrackTool extends BaseTool { async execute(params: Record): Promise { try { - this.validateParameters(params); + const normalizedParams = normalizeOptionalTrackIdParam(params); + this.validateParameters(normalizedParams); - const trackId = params.track_id as string | undefined; - const trackName = params.track_name as string | undefined; - const instrumentName = this.normalizeOptionalString(params.instrument); - const newTrackName = this.normalizeOptionalString(params.new_track_name); + const trackId = normalizedParams.track_id as string | undefined; + const trackName = normalizedParams.track_name as string | undefined; + const instrumentName = this.normalizeOptionalString(normalizedParams.instrument); + const newTrackName = this.normalizeOptionalString(normalizedParams.new_track_name); if (!trackId && !trackName) { return this.createErrorResult('Either track_id or track_name must be provided.'); diff --git a/src/agent/tools/trackIdNormalization.ts b/src/agent/tools/trackIdNormalization.ts new file mode 100644 index 0000000..2f59857 --- /dev/null +++ b/src/agent/tools/trackIdNormalization.ts @@ -0,0 +1,16 @@ +export function normalizeOptionalTrackIdParam(params: Record): Record { + const rawTrackId = params.track_id; + + if (rawTrackId === undefined || rawTrackId === null || typeof rawTrackId === 'string') { + return params; + } + + if (typeof rawTrackId === 'number' && Number.isFinite(rawTrackId) && Number.isInteger(rawTrackId)) { + return { + ...params, + track_id: String(rawTrackId), + }; + } + + return params; +}