From 0f7e6b979f46ccd9c9b89cc09d1c5098025a10e4 Mon Sep 17 00:00:00 2001 From: Eric Rakestraw Date: Sat, 23 May 2026 16:03:20 +0000 Subject: [PATCH] Deduplicate locks add/remove session-id and source parsing --- internal/app/operator_helpers_test.go | 47 ++++++++++++++++++++++ internal/app/operator_locks.go | 58 ++------------------------- internal/app/session_args.go | 34 ++++++++++++++++ 3 files changed, 85 insertions(+), 54 deletions(-) diff --git a/internal/app/operator_helpers_test.go b/internal/app/operator_helpers_test.go index 760e82c..5f313a6 100644 --- a/internal/app/operator_helpers_test.go +++ b/internal/app/operator_helpers_test.go @@ -593,6 +593,53 @@ func TestExecuteLocksRequireSessionID(t *testing.T) { } } +func TestExecuteLocksMutationRejectsSessionIDMismatch(t *testing.T) { + workspaceRoot := t.TempDir() + pipelinePath, campaignPath, sessionPath := writeValidConfigFiles(t, workspaceRoot) + fake := &storage.FakeBackend{} + var storeInitCalls int + restoreAppConfigTestGlobals(t, fake, &storeInitCalls, []string{sessionPath}) + + tests := []struct { + name string + args []string + }{ + { + name: "add mismatch", + args: []string{ + "session", "locks", "add", "2026-05-03", "narratio.transcript.final_trimmed", + "--session-id", "2026-05-04", + "--config", pipelinePath, + "--campaign-file", campaignPath, + "--session", sessionPath, + }, + }, + { + name: "remove mismatch", + args: []string{ + "session", "locks", "remove", "2026-05-03", "narratio.transcript.final_trimmed", + "--session-id", "2026-05-04", + "--config", pipelinePath, + "--campaign-file", campaignPath, + "--session", sessionPath, + }, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var stdout bytes.Buffer + var stderr bytes.Buffer + code := Execute(tt.args, &stdout, &stderr) + if code == 0 { + t.Fatal("exit code = 0, want non-zero") + } + if !strings.Contains(stderr.String(), "does not match expected session id") { + t.Fatalf("stderr = %q, want session-id mismatch guidance", stderr.String()) + } + }) + } +} + func TestExecuteLocksCannotModifyStaticLocks(t *testing.T) { workspaceRoot := t.TempDir() pipelinePath, campaignPath, sessionPath := writeValidConfigFiles(t, workspaceRoot) diff --git a/internal/app/operator_locks.go b/internal/app/operator_locks.go index cf9cb78..d9435bb 100644 --- a/internal/app/operator_locks.go +++ b/internal/app/operator_locks.go @@ -48,13 +48,6 @@ func LocksList(ctx context.Context, args []string, out io.Writer) error { // LocksAdd adds or updates one remote lock. func LocksAdd(ctx context.Context, args []string, out io.Writer) error { - var positionalSessionID string - var source string - if len(args) >= 2 && !isCLIFlagToken(args[0]) && !isCLIFlagToken(args[1]) { - positionalSessionID = strings.TrimSpace(args[0]) - source = strings.TrimSpace(args[1]) - args = append([]string(nil), args[2:]...) - } fs := flag.NewFlagSet("locks add", flag.ContinueOnError) fs.SetOutput(io.Discard) var flags commonConfigFlags @@ -63,26 +56,8 @@ func LocksAdd(ctx context.Context, args []string, out io.Writer) error { addCommonConfigFlags(fs, &flags) fs.StringVar(&reason, "reason", "", "lock reason") fs.BoolVar(&force, "force", false, "update existing remote lock") - if err := fs.Parse(args); err != nil { - return fmt.Errorf("locks add: invalid flags: %w", err) - } - if source == "" { - switch fs.NArg() { - case 2: - positionalSessionID = strings.TrimSpace(fs.Arg(0)) - source = strings.TrimSpace(fs.Arg(1)) - case 1: - if strings.TrimSpace(flags.sessionID) == "" { - return fmt.Errorf("locks add: expected session_id and source id") - } - source = strings.TrimSpace(fs.Arg(0)) - default: - return fmt.Errorf("locks add: expected session_id and source id") - } - } else if fs.NArg() != 0 { - return fmt.Errorf("locks add: unexpected positional arguments") - } - if err := applyPositionalSessionID("locks add", positionalSessionID, &flags.sessionID); err != nil { + source, err := parseSessionIDAndOnePositionalArg("locks add", "source id", fs, args, &flags.sessionID) + if err != nil { return err } if strings.TrimSpace(flags.sessionID) == "" { @@ -116,37 +91,12 @@ func LocksAdd(ctx context.Context, args []string, out io.Writer) error { // LocksRemove removes one remote lock. func LocksRemove(ctx context.Context, args []string, out io.Writer) error { - var positionalSessionID string - var source string - if len(args) >= 2 && !isCLIFlagToken(args[0]) && !isCLIFlagToken(args[1]) { - positionalSessionID = strings.TrimSpace(args[0]) - source = strings.TrimSpace(args[1]) - args = append([]string(nil), args[2:]...) - } fs := flag.NewFlagSet("locks remove", flag.ContinueOnError) fs.SetOutput(io.Discard) var flags commonConfigFlags addCommonConfigFlags(fs, &flags) - if err := fs.Parse(args); err != nil { - return fmt.Errorf("locks remove: invalid flags: %w", err) - } - if source == "" { - switch fs.NArg() { - case 2: - positionalSessionID = strings.TrimSpace(fs.Arg(0)) - source = strings.TrimSpace(fs.Arg(1)) - case 1: - if strings.TrimSpace(flags.sessionID) == "" { - return fmt.Errorf("locks remove: expected session_id and source id") - } - source = strings.TrimSpace(fs.Arg(0)) - default: - return fmt.Errorf("locks remove: expected session_id and source id") - } - } else if fs.NArg() != 0 { - return fmt.Errorf("locks remove: unexpected positional arguments") - } - if err := applyPositionalSessionID("locks remove", positionalSessionID, &flags.sessionID); err != nil { + source, err := parseSessionIDAndOnePositionalArg("locks remove", "source id", fs, args, &flags.sessionID) + if err != nil { return err } if strings.TrimSpace(flags.sessionID) == "" { diff --git a/internal/app/session_args.go b/internal/app/session_args.go index 928756f..9fd9849 100644 --- a/internal/app/session_args.go +++ b/internal/app/session_args.go @@ -59,3 +59,37 @@ func parseSessionAwareFlags(command string, fs *flag.FlagSet, args []string, ses } return resolveParsedSessionID(command, positionalSessionID, fs, sessionID) } + +func parseSessionIDAndOnePositionalArg(command, argName string, fs *flag.FlagSet, args []string, sessionID *string) (string, error) { + var positionalSessionID string + value := "" + if len(args) >= 2 && !isCLIFlagToken(args[0]) && !isCLIFlagToken(args[1]) { + positionalSessionID = strings.TrimSpace(args[0]) + value = strings.TrimSpace(args[1]) + args = append([]string(nil), args[2:]...) + } + + if err := fs.Parse(args); err != nil { + return "", fmt.Errorf("%s: invalid flags: %w", command, err) + } + if value == "" { + switch fs.NArg() { + case 2: + positionalSessionID = strings.TrimSpace(fs.Arg(0)) + value = strings.TrimSpace(fs.Arg(1)) + case 1: + if strings.TrimSpace(*sessionID) == "" { + return "", fmt.Errorf("%s: expected session_id and %s", command, argName) + } + value = strings.TrimSpace(fs.Arg(0)) + default: + return "", fmt.Errorf("%s: expected session_id and %s", command, argName) + } + } else if fs.NArg() != 0 { + return "", fmt.Errorf("%s: unexpected positional arguments", command) + } + if err := applyPositionalSessionID(command, positionalSessionID, sessionID); err != nil { + return "", err + } + return value, nil +}