Deduplicate locks add/remove session-id and source parsing
This commit is contained in:
@@ -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) {
|
func TestExecuteLocksCannotModifyStaticLocks(t *testing.T) {
|
||||||
workspaceRoot := t.TempDir()
|
workspaceRoot := t.TempDir()
|
||||||
pipelinePath, campaignPath, sessionPath := writeValidConfigFiles(t, workspaceRoot)
|
pipelinePath, campaignPath, sessionPath := writeValidConfigFiles(t, workspaceRoot)
|
||||||
|
|||||||
@@ -48,13 +48,6 @@ func LocksList(ctx context.Context, args []string, out io.Writer) error {
|
|||||||
|
|
||||||
// LocksAdd adds or updates one remote lock.
|
// LocksAdd adds or updates one remote lock.
|
||||||
func LocksAdd(ctx context.Context, args []string, out io.Writer) error {
|
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 := flag.NewFlagSet("locks add", flag.ContinueOnError)
|
||||||
fs.SetOutput(io.Discard)
|
fs.SetOutput(io.Discard)
|
||||||
var flags commonConfigFlags
|
var flags commonConfigFlags
|
||||||
@@ -63,26 +56,8 @@ func LocksAdd(ctx context.Context, args []string, out io.Writer) error {
|
|||||||
addCommonConfigFlags(fs, &flags)
|
addCommonConfigFlags(fs, &flags)
|
||||||
fs.StringVar(&reason, "reason", "", "lock reason")
|
fs.StringVar(&reason, "reason", "", "lock reason")
|
||||||
fs.BoolVar(&force, "force", false, "update existing remote lock")
|
fs.BoolVar(&force, "force", false, "update existing remote lock")
|
||||||
if err := fs.Parse(args); err != nil {
|
source, err := parseSessionIDAndOnePositionalArg("locks add", "source id", fs, args, &flags.sessionID)
|
||||||
return fmt.Errorf("locks add: invalid flags: %w", err)
|
if err != nil {
|
||||||
}
|
|
||||||
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 {
|
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if strings.TrimSpace(flags.sessionID) == "" {
|
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.
|
// LocksRemove removes one remote lock.
|
||||||
func LocksRemove(ctx context.Context, args []string, out io.Writer) error {
|
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 := flag.NewFlagSet("locks remove", flag.ContinueOnError)
|
||||||
fs.SetOutput(io.Discard)
|
fs.SetOutput(io.Discard)
|
||||||
var flags commonConfigFlags
|
var flags commonConfigFlags
|
||||||
addCommonConfigFlags(fs, &flags)
|
addCommonConfigFlags(fs, &flags)
|
||||||
if err := fs.Parse(args); err != nil {
|
source, err := parseSessionIDAndOnePositionalArg("locks remove", "source id", fs, args, &flags.sessionID)
|
||||||
return fmt.Errorf("locks remove: invalid flags: %w", err)
|
if err != nil {
|
||||||
}
|
|
||||||
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 {
|
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if strings.TrimSpace(flags.sessionID) == "" {
|
if strings.TrimSpace(flags.sessionID) == "" {
|
||||||
|
|||||||
@@ -59,3 +59,37 @@ func parseSessionAwareFlags(command string, fs *flag.FlagSet, args []string, ses
|
|||||||
}
|
}
|
||||||
return resolveParsedSessionID(command, positionalSessionID, fs, sessionID)
|
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
|
||||||
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user