diff --git a/cmd/relayfile-cli/main.go b/cmd/relayfile-cli/main.go index bcc59d09..9f0aa45b 100644 --- a/cmd/relayfile-cli/main.go +++ b/cmd/relayfile-cli/main.go @@ -6381,6 +6381,7 @@ func runMount(args []string) error { fs := flag.NewFlagSet("mount", flag.ContinueOnError) fs.SetOutput(io.Discard) + creds, _ := loadCredentials() server := fs.String("server", resolveServer("", credentials{}), "relayfile server URL") token := fs.String("token", strings.TrimSpace(os.Getenv("RELAYFILE_TOKEN")), "bearer token") credsFile := fs.String("creds-file", strings.TrimSpace(os.Getenv("RELAYFILE_MOUNT_CREDS_FILE")), "delegated relayfile credentials file") @@ -6459,8 +6460,11 @@ func runMount(args []string) error { stateFileProvided := false stateDirProvided := false mountKindProvided := false + serverProvided := false fs.Visit(func(parsed *flag.Flag) { switch parsed.Name { + case "server": + serverProvided = true case "local-layout": localLayoutProvided = true case "state-file": @@ -6487,6 +6491,7 @@ func runMount(args []string) error { canonicalWorkspaceID := "" requestedWorkspace := "" delegatedCredsPath := resolveDelegatedCredentialsPath(*credsFile) + _, delegatedCredsExplicit := explicitDelegatedCredentialsPath(*credsFile) usesDelegatedWorkspace := false initialCredExpiresAt := "" if fs.NArg() > 0 { @@ -6495,25 +6500,47 @@ func runMount(args []string) error { if tokenValue == "" { bundle, path, berr := loadDelegatedCredentialsForRequest(*credsFile, requestedWorkspace, defaultJoinScopes) if berr != nil { - return fmt.Errorf("resolve delegated relayfile credentials: %w", berr) - } - delegatedCredsPath = path - bundle, berr = refreshDelegatedCredentials(path, bundle, false) - if berr != nil { - return fmt.Errorf("refresh delegated relayfile credentials: %w", berr) - } - canonicalWorkspaceID = bundle.Workspace() - if requestedWorkspace != "" && !workspaceRequestMatchesDelegatedCredentials(requestedWorkspace, canonicalWorkspaceID) { - return fmt.Errorf( - "relayfile mount without --token uses delegated relayfile workspace %s; pass --token for explicit workspace %q or re-bootstrap delegated credentials for that workspace", - canonicalWorkspaceID, - requestedWorkspace, - ) + if delegatedCredsExplicit { + return fmt.Errorf("resolve delegated relayfile credentials: %w", berr) + } + tokenValue = strings.TrimSpace(creds.Token) + if tokenValue == "" { + return fmt.Errorf("resolve delegated relayfile credentials: %w", berr) + } + if !serverProvided { + *server = resolveServer("", creds) + } + delegatedCredsPath = "" + } else { + delegatedCredsPath = path + bundle, berr = refreshDelegatedCredentials(path, bundle, false) + if berr != nil { + if delegatedCredsExplicit { + return fmt.Errorf("refresh delegated relayfile credentials: %w", berr) + } + tokenValue = strings.TrimSpace(creds.Token) + if tokenValue == "" { + return fmt.Errorf("refresh delegated relayfile credentials: %w", berr) + } + if !serverProvided { + *server = resolveServer("", creds) + } + delegatedCredsPath = "" + } else { + canonicalWorkspaceID = bundle.Workspace() + if requestedWorkspace != "" && !workspaceRequestMatchesDelegatedCredentials(requestedWorkspace, canonicalWorkspaceID) { + return fmt.Errorf( + "relayfile mount without --token uses delegated relayfile workspace %s; pass --token for explicit workspace %q or re-bootstrap delegated credentials for that workspace", + canonicalWorkspaceID, + requestedWorkspace, + ) + } + tokenValue = bundle.BearerToken() + initialCredExpiresAt = bundle.BearerExpiresAt() + *server = strings.TrimRight(bundle.ServerURL(), "/") + usesDelegatedWorkspace = true + } } - tokenValue = bundle.BearerToken() - initialCredExpiresAt = bundle.BearerExpiresAt() - *server = strings.TrimRight(bundle.ServerURL(), "/") - usesDelegatedWorkspace = true } var err error switch fs.NArg() { @@ -7099,20 +7126,43 @@ func prepareWorkspaceCommandClient(workspaceValue, serverFlag, tokenFlag string, tokenValue := resolveExplicitToken(tokenFlag) directToken := tokenValue != "" credsFile := "" + _, delegatedCredsExplicit := explicitDelegatedCredentialsPath("") var bundle delegatedauth.Bundle var err error if !directToken && strings.TrimSpace(tokenValue) == "" { - bundle, credsFile, err = loadOrBootstrapDelegatedCredentials(workspaceValue, requestedScopes) - if err != nil { - return nil, fmt.Errorf("resolve delegated relayfile credentials: %w", err) + if delegatedCredsExplicit { + bundle, credsFile, err = loadDelegatedCredentialsForRequest("", workspaceValue, requestedScopes) + } else { + bundle, credsFile, err = loadOrBootstrapDelegatedCredentials(workspaceValue, requestedScopes) } - bundle, err = refreshDelegatedCredentials(credsFile, bundle, false) if err != nil { - return nil, fmt.Errorf("refresh delegated relayfile credentials: %w", err) - } - tokenValue = bundle.BearerToken() - if strings.TrimSpace(serverFlag) == "" { - serverFlag = bundle.ServerURL() + if delegatedCredsExplicit { + return nil, fmt.Errorf("resolve delegated relayfile credentials: %w", err) + } + tokenValue = strings.TrimSpace(creds.Token) + if tokenValue == "" { + return nil, fmt.Errorf("resolve delegated relayfile credentials: %w", err) + } + directToken = true + credsFile = "" + } else { + bundle, err = refreshDelegatedCredentials(credsFile, bundle, false) + if err != nil { + if delegatedCredsExplicit { + return nil, fmt.Errorf("refresh delegated relayfile credentials: %w", err) + } + tokenValue = strings.TrimSpace(creds.Token) + if tokenValue == "" { + return nil, fmt.Errorf("refresh delegated relayfile credentials: %w", err) + } + directToken = true + credsFile = "" + } else { + tokenValue = bundle.BearerToken() + if strings.TrimSpace(serverFlag) == "" { + serverFlag = bundle.ServerURL() + } + } } } workspaceID := "" @@ -7181,10 +7231,10 @@ func prepareWorkspaceCommandClient(workspaceValue, serverFlag, tokenFlag string, } func workspaceRecordForCommand(workspaceValue, workspaceID string) workspaceRecord { - if record, ok := workspaceRecordByName(strings.TrimSpace(workspaceValue)); ok { + if record, ok := workspaceRecordByID(workspaceID); ok { return normalizeWorkspaceCommandRecord(record, workspaceID) } - if record, ok := workspaceRecordByID(workspaceID); ok { + if record, ok := workspaceRecordByName(strings.TrimSpace(workspaceValue)); ok { return normalizeWorkspaceCommandRecord(record, workspaceID) } return workspaceRecord{ @@ -10591,21 +10641,27 @@ func resolveWorkspaceRecord(nameOrID string) (workspaceRecord, error) { func resolveWorkspaceIDWithToken(value, token string) (string, error) { value = strings.TrimSpace(value) if value != "" { - if id, ok := catalogWorkspaceID(value); ok { + if id, ok, err := catalogWorkspaceIDForRequest(value, token); err != nil { + return "", err + } else if ok { return id, nil } return value, nil } if workspaceID := strings.TrimSpace(os.Getenv("RELAYFILE_WORKSPACE")); workspaceID != "" { - if id, ok := catalogWorkspaceID(workspaceID); ok { + if id, ok, err := catalogWorkspaceIDForRequest(workspaceID, token); err != nil { + return "", err + } else if ok { return id, nil } return workspaceID, nil } if workspaceID := workspaceIDFromToken(token); workspaceID != "" { - if id, ok := catalogWorkspaceID(workspaceID); ok { + if id, ok, err := catalogWorkspaceIDForRequest(workspaceID, token); err != nil { + return "", err + } else if ok { return id, nil } return workspaceID, nil @@ -10621,7 +10677,9 @@ func resolveWorkspaceIDWithToken(value, token string) (string, error) { } defaultName := strings.TrimSpace(catalog.Default) if defaultName != "" { - if id, ok := catalogWorkspaceIDFromCatalog(catalog, defaultName); ok { + if id, ok, err := catalogWorkspaceIDFromCatalogForRequest(catalog, defaultName, token); err != nil { + return "", err + } else if ok { return id, nil } return defaultName, nil @@ -10629,6 +10687,72 @@ func resolveWorkspaceIDWithToken(value, token string) (string, error) { return "", errors.New("workspace is required; pass WORKSPACE, set RELAYFILE_WORKSPACE, or run 'agent-relay workspace switch NAME'") } +func catalogWorkspaceIDForRequest(name, token string) (string, bool, error) { + catalog, err := loadWorkspaceCatalog() + if err != nil { + return "", false, nil + } + return catalogWorkspaceIDFromCatalogForRequest(catalog, name, token) +} + +func catalogWorkspaceIDFromCatalogForRequest(catalog workspaceCatalog, name, token string) (string, bool, error) { + name = strings.TrimSpace(name) + if name == "" { + return "", false, nil + } + + // An exact ID is already unambiguous, even if another record happens to + // reuse that value as its display name. + for _, workspace := range catalog.Workspaces { + if strings.TrimSpace(workspace.ID) == name { + return name, true, nil + } + } + + matches := make([]workspaceRecord, 0, 1) + ids := map[string]struct{}{} + for _, workspace := range catalog.Workspaces { + if strings.TrimSpace(workspace.Name) != name { + continue + } + id := strings.TrimSpace(workspace.ID) + if id == "" { + id = name + } + matches = append(matches, workspace) + ids[id] = struct{}{} + } + if len(ids) == 0 { + return "", false, nil + } + if len(ids) == 1 { + for id := range ids { + return id, true, nil + } + } + + tokenWorkspaceID := workspaceIDFromToken(token) + if tokenWorkspaceID != "" { + matchedID := "" + for _, workspace := range matches { + if strings.TrimSpace(workspace.ID) != tokenWorkspaceID && strings.TrimSpace(workspace.RelayWorkspaceID) != tokenWorkspaceID { + continue + } + if matchedID != "" { + return "", false, fmt.Errorf("workspace %q is ambiguous in %s; pass an exact workspace id", name, workspacesPath()) + } + matchedID = strings.TrimSpace(workspace.ID) + if matchedID == "" { + matchedID = name + } + } + if matchedID != "" { + return matchedID, true, nil + } + } + return "", false, fmt.Errorf("workspace %q is ambiguous in %s; pass an exact workspace id", name, workspacesPath()) +} + func catalogWorkspaceID(name string) (string, bool) { catalog, err := loadWorkspaceCatalog() if err != nil { @@ -11290,7 +11414,7 @@ func runningMountDaemons(localDir, workspaceID, workspaceName string) ([]mountDa seen[process.PID] = struct{}{} } - if pid, verified, strong := verifyDaemonProcessForDiscovery(localDir, workspaceID); pid != 0 { + if pid, verified, strong := verifyDaemonProcessForDiscovery(localDir, workspaceID); pid != 0 && pid != os.Getpid() { _, foundByScan := seen[pid] switch { case !processAlive(pid): @@ -11354,6 +11478,11 @@ func mountDaemonCommandMatches(command, localDir, workspaceID, workspaceName str if !commandHasMountSubcommand(fields) || commandHasOnceFlag(fields) { return false } + // `mount --background` is a transient launcher. Only the child carrying + // `--daemonized` serves the mount and should participate in discovery. + if commandHasEnabledBoolFlag(fields, "background") && !commandHasEnabledBoolFlag(fields, "daemonized") { + return false + } targets := daemonWorkspaceTargets(workspaceID, workspaceName) if commandMatchesWorkspace(fields, targets) { return true @@ -11386,12 +11515,18 @@ func commandHasMountSubcommand(fields []string) bool { } func commandHasOnceFlag(fields []string) bool { + return commandHasEnabledBoolFlag(fields, "once") +} + +func commandHasEnabledBoolFlag(fields []string, name string) bool { + longFlag := "--" + name + shortFlag := "-" + name for _, field := range fields { - if field == "--once" || field == "-once" { + if field == longFlag || field == shortFlag { return true } - if strings.HasPrefix(field, "--once=") || strings.HasPrefix(field, "-once=") { - value := strings.TrimSpace(strings.TrimPrefix(strings.TrimPrefix(field, "--once="), "-once=")) + if strings.HasPrefix(field, longFlag+"=") || strings.HasPrefix(field, shortFlag+"=") { + value := strings.TrimSpace(strings.TrimPrefix(strings.TrimPrefix(field, longFlag+"="), shortFlag+"=")) if value == "" { return true } diff --git a/cmd/relayfile-cli/main_test.go b/cmd/relayfile-cli/main_test.go index bb9732ce..71b4e9c2 100644 --- a/cmd/relayfile-cli/main_test.go +++ b/cmd/relayfile-cli/main_test.go @@ -11,7 +11,9 @@ import ( "net/http/httptest" "net/url" "os" + "os/exec" "path/filepath" + "runtime" "strconv" "strings" "sync/atomic" @@ -23,6 +25,19 @@ import ( "github.com/agentworkforce/relayfile/internal/mountsync" ) +const relayfileCLITestSubprocessEnv = "RELAYFILE_CLI_TEST_SUBPROCESS" + +func TestMain(m *testing.M) { + if os.Getenv(relayfileCLITestSubprocessEnv) == "1" { + if err := run(os.Args[1:], os.Stdin, os.Stdout, os.Stderr); err != nil { + fmt.Fprintln(os.Stderr, "error:", err) + os.Exit(1) + } + os.Exit(0) + } + os.Exit(m.Run()) +} + func TestWorkspaceCreateStoresCatalogEntry(t *testing.T) { t.Setenv("HOME", t.TempDir()) clearRelayfileEnv(t) @@ -2229,6 +2244,15 @@ func TestStatusSurfacesOrphanDaemonFromProcessScan(t *testing.T) { func TestMountDaemonCommandMatchesAliasesAndLocalDirBoundaries(t *testing.T) { localDir := filepath.Join(t.TempDir(), "ws") + if mountDaemonCommandMatches("relayfile mount demo "+localDir+" --background", localDir, "ws_demo", "demo") { + t.Fatalf("transient --background launcher must not match a serving daemon") + } + if !mountDaemonCommandMatches("relayfile mount demo "+localDir+" --background=false", localDir, "ws_demo", "demo") { + t.Fatalf("foreground mount with --background=false must still match a serving daemon") + } + if !mountDaemonCommandMatches("relayfile mount demo "+localDir+" --background --daemonized", localDir, "ws_demo", "demo") { + t.Fatalf("daemonized child must still match even if inherited argv contains --background") + } if !mountDaemonCommandMatches("relayfile start demo "+localDir+" --daemonized", localDir, "ws_demo", "demo") { t.Fatalf("expected relayfile start alias to match daemon command") } @@ -2243,6 +2267,32 @@ func TestMountDaemonCommandMatchesAliasesAndLocalDirBoundaries(t *testing.T) { } } +func TestRunningMountDaemonsExcludesCallerPIDFromPIDFile(t *testing.T) { + localDir := t.TempDir() + if err := ensureMirrorLayout(localDir); err != nil { + t.Fatalf("ensureMirrorLayout failed: %v", err) + } + if err := writeDaemonPIDState(mountPIDFile(localDir), daemonPIDState{ + PID: os.Getpid(), + LocalDir: localDir, + StartedAt: time.Now().UTC().Format(time.RFC3339), + Executable: resolvedSelfExecutable(), + }); err != nil { + t.Fatalf("writeDaemonPIDState failed: %v", err) + } + oldList := listProcessCommands + listProcessCommands = func() ([]processCommandSnapshot, error) { return nil, nil } + t.Cleanup(func() { listProcessCommands = oldList }) + + running, stalePID, err := runningMountDaemons(localDir, "ws_demo", "demo") + if err != nil { + t.Fatalf("runningMountDaemons failed: %v", err) + } + if len(running) != 0 || stalePID != 0 { + t.Fatalf("caller pid reported as competing daemon: running=%+v stalePID=%d", running, stalePID) + } +} + func TestRestartRequiresRecordedLocalMirror(t *testing.T) { t.Setenv("HOME", t.TempDir()) clearRelayfileEnv(t) @@ -2622,6 +2672,74 @@ func TestPrepareBackgroundMountLayoutPreservesScopedArtifactAbsence(t *testing.T } } +func TestSpawnBackgroundMountProcessRegistersRealChild(t *testing.T) { + t.Setenv("HOME", t.TempDir()) + clearRelayfileEnv(t) + t.Setenv(relayfileCLITestSubprocessEnv, "1") + + token := testJWTWithWorkspace("ws_background") + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if got, want := r.Header.Get("Authorization"), "Bearer "+token; got != want { + t.Fatalf("Authorization = %q, want %q", got, want) + } + w.Header().Set("Content-Type", "application/json") + switch r.URL.Path { + case "/v1/workspaces/ws_background/fs/export": + _, _ = w.Write([]byte(`[]`)) + case "/v1/workspaces/ws_background/fs/events": + _, _ = w.Write([]byte(`{"events":[]}`)) + case "/v1/workspaces/ws_background/sync/status": + _, _ = w.Write([]byte(`{"workspaceId":"ws_background","providers":[]}`)) + default: + t.Fatalf("unexpected request: %s %s", r.Method, r.URL.Path) + } + })) + defer server.Close() + + localDir := filepath.Join(t.TempDir(), "mirror") + stateDir := t.TempDir() + pidFile := mountPIDFile(localDir) + logFile := mountLogFile(localDir) + args := []string{ + "ws_background", localDir, + "--server", server.URL, + "--token", token, + "--state-dir", stateDir, + "--interval", "1h", + "--websocket=false", + } + setPathWithPSOnly(t) + if err := spawnBackgroundMountProcess(args, []string{"/"}, localDir, pidFile, logFile, mountscope.LayoutExact); err != nil { + logBytes, _ := os.ReadFile(logFile) + t.Fatalf("spawnBackgroundMountProcess failed: %v\nlog:\n%s", err, logBytes) + } + state, structured := readDaemonPIDStateFile(pidFile) + t.Cleanup(func() { + if state.PID > 0 && processAlive(state.PID) { + if process, err := os.FindProcess(state.PID); err == nil { + _ = forceDaemonStop(process) + } + } + }) + if !structured || !state.Registered || state.PID <= 0 { + t.Fatalf("background child did not register: structured=%v state=%+v", structured, state) + } + process, err := os.FindProcess(state.PID) + if err != nil { + t.Fatalf("find background child: %v", err) + } + if err := signalDaemonStop(process); err != nil && !isProcessAlreadyGone(err) { + t.Fatalf("stop background child: %v", err) + } + deadline := time.Now().Add(5 * time.Second) + for processAlive(state.PID) && time.Now().Before(deadline) { + time.Sleep(25 * time.Millisecond) + } + if processAlive(state.PID) { + t.Fatalf("background child %d did not exit after stop", state.PID) + } +} + func TestPrepareScopedCatalogRootRemovesOnlyGeneratedArtifacts(t *testing.T) { localRoot := filepath.Join(t.TempDir(), "mirror") if err := ensureMirrorLayout(localRoot); err != nil { @@ -3800,6 +3918,7 @@ func TestMountSkipsDataPlaneWhenDelegatedRefreshRejected(t *testing.T) { refreshCalls := 0 dataPlaneCalls := 0 + savedToken := testJWTWithWorkspace("ws_refresh") var server *httptest.Server server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "application/json") @@ -3810,7 +3929,17 @@ func TestMountSkipsDataPlaneWhenDelegatedRefreshRejected(t *testing.T) { _, _ = w.Write([]byte(`{"code":"delegation_expired","message":"delegation expired"}`)) case "/v1/workspaces/ws_refresh/fs/events", "/v1/workspaces/ws_refresh/fs/export", "/v1/workspaces/ws_refresh/sync/status": dataPlaneCalls++ - t.Fatalf("data-plane call should be skipped after rejected refresh: %s", r.URL.Path) + if got, want := r.Header.Get("Authorization"), "Bearer "+savedToken; got != want { + t.Fatalf("fallback Authorization = %q, want %q", got, want) + } + switch r.URL.Path { + case "/v1/workspaces/ws_refresh/fs/export": + _, _ = w.Write([]byte(`[]`)) + case "/v1/workspaces/ws_refresh/fs/events": + _, _ = w.Write([]byte(`{"events":[]}`)) + default: + _, _ = w.Write([]byte(`{"workspaceId":"ws_refresh","providers":[]}`)) + } default: t.Fatalf("unexpected path: %s", r.URL.Path) } @@ -3825,6 +3954,9 @@ func TestMountSkipsDataPlaneWhenDelegatedRefreshRejected(t *testing.T) { RefreshTokenExpiresAt: time.Now().Add(24 * time.Hour).UTC().Format(time.RFC3339), RelayauthURL: server.URL, }) + if err := saveCredentials(credentials{Server: server.URL, Token: savedToken}); err != nil { + t.Fatalf("save fallback login credentials: %v", err) + } before, err := os.ReadFile(credsPath) if err != nil { t.Fatalf("read delegated credentials before mount failed: %v", err) @@ -3833,6 +3965,7 @@ func TestMountSkipsDataPlaneWhenDelegatedRefreshRejected(t *testing.T) { err = run([]string{ "mount", "ws_refresh", localDir, "--server", server.URL, + "--creds-file", credsPath, "--once", "--websocket=false", }, strings.NewReader(""), &bytes.Buffer{}, &bytes.Buffer{}) @@ -3873,6 +4006,25 @@ func clearRelayfileEnv(t *testing.T) { t.Setenv("AGENT_RELAY_BIN", "") } +func setPathWithPSOnly(t *testing.T) { + t.Helper() + if runtime.GOOS == "windows" { + t.Skip("requires Unix ps-based mount process discovery") + } + psPath, err := exec.LookPath("ps") + if err != nil { + t.Skipf("ps is unavailable for mount process discovery: %v", err) + } + binDir := t.TempDir() + if err := os.Symlink(psPath, filepath.Join(binDir, "ps")); err != nil { + t.Fatalf("link ps into isolated PATH: %v", err) + } + t.Setenv("PATH", binDir) + if _, err := exec.LookPath("agent-relay"); err == nil { + t.Fatal("isolated PATH unexpectedly contains agent-relay") + } +} + func installFakeAgentRelay(t *testing.T, scriptBody string) string { t.Helper() path := filepath.Join(t.TempDir(), "agent-relay") @@ -6318,6 +6470,201 @@ func TestLoginWithExplicitTokenPersistsServerCreds(t *testing.T) { } } +func TestLoginCredentialsAuthorizeMountAndStatusWithoutAgentRelay(t *testing.T) { + t.Setenv("HOME", t.TempDir()) + clearRelayfileEnv(t) + setPathWithPSOnly(t) + + localDir := t.TempDir() + token := testJWTWithWorkspace("ws_saved") + var healthCalls, exportCalls, eventsCalls, statusCalls int + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if got, want := r.Header.Get("Authorization"), "Bearer "+token; got != want { + t.Fatalf("Authorization = %q, want %q", got, want) + } + w.Header().Set("Content-Type", "application/json") + switch r.URL.Path { + case "/health": + healthCalls++ + w.WriteHeader(http.StatusOK) + case "/v1/workspaces/ws_saved/fs/export": + exportCalls++ + _, _ = w.Write([]byte(`[]`)) + case "/v1/workspaces/ws_saved/fs/events": + eventsCalls++ + _, _ = w.Write([]byte(`{"events":[]}`)) + case "/v1/workspaces/ws_saved/sync/status": + statusCalls++ + _, _ = w.Write([]byte(`{"workspaceId":"ws_saved","providers":[]}`)) + default: + t.Fatalf("unexpected request: %s %s", r.Method, r.URL.Path) + } + })) + defer server.Close() + + now := time.Now().UTC().Format(time.RFC3339) + if err := saveWorkspaceCatalog(workspaceCatalog{ + Default: "demo", + Workspaces: []workspaceRecord{{ + Name: "demo", + ID: "ws_saved", + CreatedAt: now, + }}, + }); err != nil { + t.Fatalf("saveWorkspaceCatalog failed: %v", err) + } + + var stdout bytes.Buffer + if err := run([]string{"login", "--server", server.URL, "--token", token}, strings.NewReader(""), &stdout, &stdout); err != nil { + t.Fatalf("run login failed: %v\noutput:\n%s", err, stdout.String()) + } + stdout.Reset() + if err := run([]string{"mount", "demo", localDir, "--once", "--websocket=false"}, strings.NewReader(""), &stdout, &stdout); err != nil { + t.Fatalf("run mount with saved login failed: %v\noutput:\n%s", err, stdout.String()) + } + stdout.Reset() + if err := run([]string{"status", "demo"}, strings.NewReader(""), &stdout, &stdout); err != nil { + t.Fatalf("run status with saved login failed: %v\noutput:\n%s", err, stdout.String()) + } + if healthCalls != 1 || exportCalls == 0 || eventsCalls == 0 || statusCalls == 0 { + t.Fatalf("unexpected request counts: health=%d export=%d events=%d status=%d", healthCalls, exportCalls, eventsCalls, statusCalls) + } + creds, err := loadCredentials() + if err != nil { + t.Fatalf("loadCredentials after mount/status failed: %v", err) + } + if creds.Server != server.URL || creds.Token != token { + t.Fatalf("saved login changed after mount/status: %#v", creds) + } +} + +func TestExplicitDelegatedCredentialsDoNotFallBackToSavedLogin(t *testing.T) { + tests := []struct { + name string + command string + configure func(t *testing.T, path string) []string + }{ + { + name: "mount flag selects missing file", + command: "mount", + configure: func(t *testing.T, path string) []string { + return []string{"--creds-file", path} + }, + }, + { + name: "mount environment selects malformed file", + command: "mount", + configure: func(t *testing.T, path string) []string { + if err := os.WriteFile(path, []byte("not-json"), 0o600); err != nil { + t.Fatalf("write malformed delegated credentials: %v", err) + } + t.Setenv("RELAYFILE_MOUNT_CREDS_FILE", path) + return nil + }, + }, + { + name: "status environment selects malformed file", + command: "status", + configure: func(t *testing.T, path string) []string { + if err := os.WriteFile(path, []byte("not-json"), 0o600); err != nil { + t.Fatalf("write malformed delegated credentials: %v", err) + } + t.Setenv("RELAYFILE_MOUNT_CREDS_FILE", path) + return nil + }, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + t.Setenv("HOME", t.TempDir()) + clearRelayfileEnv(t) + + token := testJWTWithWorkspace("ws_saved") + requestCount := 0 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + requestCount++ + w.Header().Set("Content-Type", "application/json") + switch r.URL.Path { + case "/v1/workspaces/ws_saved/fs/export": + _, _ = w.Write([]byte(`[]`)) + case "/v1/workspaces/ws_saved/fs/events": + _, _ = w.Write([]byte(`{"events":[]}`)) + case "/v1/workspaces/ws_saved/sync/status": + _, _ = w.Write([]byte(`{"workspaceId":"ws_saved","providers":[]}`)) + default: + t.Fatalf("unexpected fallback request: %s %s", r.Method, r.URL.Path) + } + })) + defer server.Close() + if err := saveCredentials(credentials{Server: server.URL, Token: token}); err != nil { + t.Fatalf("save credentials: %v", err) + } + + selectedPath := filepath.Join(t.TempDir(), "explicit-delegated.json") + extraArgs := tc.configure(t, selectedPath) + args := []string{tc.command, "ws_saved"} + if tc.command == "mount" { + args = append(args, t.TempDir(), "--once", "--websocket=false") + } + args = append(args, extraArgs...) + err := run(args, strings.NewReader(""), &bytes.Buffer{}, &bytes.Buffer{}) + if err == nil || !strings.Contains(err.Error(), "resolve delegated relayfile credentials") { + t.Fatalf("explicit delegated credential failure = %v, want resolution error", err) + } + if requestCount != 0 { + t.Fatalf("saved login was used after explicit delegated credential failure: %d request(s)", requestCount) + } + }) + } +} + +func TestStatusUsesSavedTokenToDisambiguateDuplicateWorkspaceNames(t *testing.T) { + t.Setenv("HOME", t.TempDir()) + clearRelayfileEnv(t) + setPathWithPSOnly(t) + + now := time.Now().UTC().Format(time.RFC3339) + if err := saveWorkspaceCatalog(workspaceCatalog{ + Default: "demo", + Workspaces: []workspaceRecord{ + {Name: "demo", ID: "ws_stale", RelayWorkspaceID: "relay_ws_stale", LocalDir: filepath.Join(t.TempDir(), "stale"), CreatedAt: now}, + {Name: "demo", ID: "ws_current", RelayWorkspaceID: "relay_ws_current", LocalDir: filepath.Join(t.TempDir(), "current"), CreatedAt: now}, + }, + }); err != nil { + t.Fatalf("saveWorkspaceCatalog failed: %v", err) + } + if _, err := resolveWorkspaceIDWithToken("demo", "opaque-token"); err == nil || !strings.Contains(err.Error(), "ambiguous") { + t.Fatalf("opaque token should fail closed on duplicate workspace names, got %v", err) + } + + // The token claim identifies the relay workspace, while API requests must + // use the matching catalog record's canonical Relayfile workspace ID. + token := testJWTWithWorkspace("relay_ws_current") + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/v1/workspaces/ws_current/sync/status" { + t.Fatalf("unexpected path: %s", r.URL.Path) + } + if got, want := r.Header.Get("Authorization"), "Bearer "+token; got != want { + t.Fatalf("Authorization = %q, want %q", got, want) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"workspaceId":"ws_current","providers":[]}`)) + })) + defer server.Close() + if err := saveCredentials(credentials{Server: server.URL, Token: token}); err != nil { + t.Fatalf("saveCredentials failed: %v", err) + } + + var stdout bytes.Buffer + if err := run([]string{"status", "demo"}, strings.NewReader(""), &stdout, &stdout); err != nil { + t.Fatalf("run status failed: %v\noutput:\n%s", err, stdout.String()) + } + if got := stdout.String(); !strings.Contains(got, "workspace ws_current (demo)") { + t.Fatalf("status used the wrong duplicate workspace record: %q", got) + } +} + // TestLoginDelegatesToAgentRelay covers the unified auth behavior: relayfile // login no longer writes its own cloud credential store; it delegates to the // canonical agent-relay login command. diff --git a/internal/mountsync/assessment_propagation_test.go b/internal/mountsync/assessment_propagation_test.go index dc1a7e08..44da593a 100644 --- a/internal/mountsync/assessment_propagation_test.go +++ b/internal/mountsync/assessment_propagation_test.go @@ -37,6 +37,28 @@ func newAssessStore(t *testing.T) *relayfile.Store { return store } +type blockingReadFileClient struct { + RemoteClient + started chan struct{} + release chan struct{} + once sync.Once +} + +func (c *blockingReadFileClient) ReadFile(ctx context.Context, workspaceID, path string) (RemoteFile, error) { + file, err := c.RemoteClient.ReadFile(ctx, workspaceID, path) + c.once.Do(func() { + close(c.started) + select { + case <-c.release: + case <-ctx.Done(): + } + }) + if err == nil && ctx.Err() != nil { + return RemoteFile{}, ctx.Err() + } + return file, err +} + // TestAssessPropagationLatencyWebSocket measures wall-clock time from a // remote write landing on the server to sandbox B's local mirror reflecting // it, with WebSocket push enabled (the default: RELAYFILE_MOUNT_WEBSOCKET @@ -334,6 +356,241 @@ func TestAssessSameFileNearSimultaneousWrite(t *testing.T) { } } +// TestAssessSameFileWebSocketBeatsDebouncedWatcher covers the real-daemon +// ordering that TestAssessSameFileNearSimultaneousWrite cannot exercise: +// fsnotify waits 100ms for a local save to settle, while a competing writer's +// WebSocket update may be applied immediately. The remote update must preserve +// the on-disk edit even though HandleLocalChange has not marked it Dirty yet. +func TestAssessSameFileWebSocketBeatsDebouncedWatcher(t *testing.T) { + t.Parallel() + store := newAssessStore(t) + workspaceID := "ws_assess_same_file_ws_debounce" + handler := newMountsyncAPIHandler(t, store) + api := httptest.NewServer(handler) + defer api.Close() + + tokenA := mustMountsyncTestJWT(t, "dev-secret", workspaceID, "SandboxA", []string{"fs:read", "fs:write"}, time.Now().Add(time.Hour)) + tokenB := mustMountsyncTestJWT(t, "dev-secret", workspaceID, "SandboxB", []string{"fs:read", "fs:write"}, time.Now().Add(time.Hour)) + const remotePath = "/notion/Shared.md" + writeMountsyncRemoteFile(t, api.Client(), api.URL, tokenA, workspaceID, remotePath, "0", "base content") + + localDirA := t.TempDir() + localDirB := t.TempDir() + wsDisabled := false + syncerA, err := NewSyncer(NewHTTPClient(api.URL, tokenA, api.Client()), SyncerOptions{ + WorkspaceID: workspaceID, RemoteRoot: "/notion", LocalRoot: localDirA, WebSocket: &wsDisabled, + }) + if err != nil { + t.Fatalf("new syncer A: %v", err) + } + syncerB, err := NewSyncer(NewHTTPClient(api.URL, tokenB, api.Client()), SyncerOptions{ + WorkspaceID: workspaceID, RemoteRoot: "/notion", LocalRoot: localDirB, WebSocket: &wsDisabled, + }) + if err != nil { + t.Fatalf("new syncer B: %v", err) + } + + ctx := context.Background() + if err := syncerA.SyncOnce(ctx); err != nil { + t.Fatalf("A bootstrap: %v", err) + } + if err := syncerB.SyncOnce(ctx); err != nil { + t.Fatalf("B bootstrap: %v", err) + } + // Revision identifiers are opaque by contract. Make B's base sort after + // the incoming store revision so any implementation that incorrectly uses + // revision ordering to gate preservation fails this regression. + trackedB := syncerB.state.Files[remotePath] + trackedB.Revision = "zzzz_opaque_base" + syncerB.state.Files[remotePath] = trackedB + localPathA := filepath.Join(localDirA, "Shared.md") + localPathB := filepath.Join(localDirB, "Shared.md") + if err := os.WriteFile(localPathA, []byte("edit from A"), 0o644); err != nil { + t.Fatalf("local write A: %v", err) + } + if err := os.WriteFile(localPathB, []byte("edit from B"), 0o644); err != nil { + t.Fatalf("local write B: %v", err) + } + + // A's watcher callback wins the server race. On B, model the WebSocket + // callback arriving before fsnotify's 100ms debounced callback. + if err := syncerA.HandleLocalChange(ctx, "Shared.md", 0); err != nil { + t.Fatalf("A push: %v", err) + } + if err := syncerB.applyWebSocketEvent(ctx, websocketEvent{ + Type: "file.updated", + Path: remotePath, + Timestamp: time.Now().UTC().Format(time.RFC3339Nano), + }); err != nil { + t.Fatalf("B websocket apply: %v", err) + } + if err := syncerB.HandleLocalChange(ctx, "Shared.md", 0); err != nil { + t.Fatalf("B delayed watcher callback: %v", err) + } + + assertLocalFileContent(t, localPathB, "edit from A") + conflicts := listConflictArtifacts(t, localDirB) + if len(conflicts) != 1 { + t.Fatalf("B should preserve its debounced local edit in one conflict artifact, got %v", conflicts) + } + artifact, err := os.ReadFile(filepath.Join(localDirB, ".relay", "conflicts", conflicts[0])) + if err != nil { + t.Fatalf("read B conflict artifact: %v", err) + } + if got := string(artifact); got != "edit from B" { + t.Fatalf("B conflict artifact = %q, want %q", got, "edit from B") + } +} + +// A duplicate or stale WebSocket event can arrive during the same debounce +// window even when the server content has not advanced from the tracked base. +// It must leave the local edit for the watcher to push without manufacturing a +// conflict, because only one side has actually changed. +func TestAssessDuplicateRemoteBasePreservesDebouncedLocalEdit(t *testing.T) { + t.Parallel() + store := newAssessStore(t) + workspaceID := "ws_assess_duplicate_remote_base" + handler := newMountsyncAPIHandler(t, store) + api := httptest.NewServer(handler) + defer api.Close() + + token := mustMountsyncTestJWT(t, "dev-secret", workspaceID, "SandboxB", []string{"fs:read", "fs:write"}, time.Now().Add(time.Hour)) + const remotePath = "/notion/Shared.md" + writeMountsyncRemoteFile(t, api.Client(), api.URL, token, workspaceID, remotePath, "0", "base content") + + localDir := t.TempDir() + wsDisabled := false + syncer, err := NewSyncer(NewHTTPClient(api.URL, token, api.Client()), SyncerOptions{ + WorkspaceID: workspaceID, RemoteRoot: "/notion", LocalRoot: localDir, WebSocket: &wsDisabled, + }) + if err != nil { + t.Fatalf("new syncer: %v", err) + } + ctx := context.Background() + if err := syncer.SyncOnce(ctx); err != nil { + t.Fatalf("bootstrap: %v", err) + } + + localPath := filepath.Join(localDir, "Shared.md") + if err := os.WriteFile(localPath, []byte("edit from B"), 0o644); err != nil { + t.Fatalf("local write: %v", err) + } + if err := syncer.applyWebSocketEvent(ctx, websocketEvent{ + Type: "file.updated", + Path: remotePath, + Timestamp: time.Now().UTC().Format(time.RFC3339Nano), + }); err != nil { + t.Fatalf("duplicate websocket apply: %v", err) + } + + assertLocalFileContent(t, localPath, "edit from B") + if conflicts := listConflictArtifacts(t, localDir); len(conflicts) != 0 { + t.Fatalf("duplicate remote base should not create a conflict artifact, got %v", conflicts) + } + if err := syncer.HandleLocalChange(ctx, "Shared.md", 0); err != nil { + t.Fatalf("delayed watcher callback: %v", err) + } + assertLocalFileContent(t, localPath, "edit from B") + remoteFile, err := syncer.client.ReadFile(ctx, workspaceID, remotePath) + if err != nil { + t.Fatalf("read server file after delayed watcher: %v", err) + } + if got := remoteFile.Content; got != "edit from B" { + t.Fatalf("server content after delayed watcher = %q, want %q", got, "edit from B") + } + if conflicts := listConflictArtifacts(t, localDir); len(conflicts) != 0 { + t.Fatalf("single-sided local edit should remain conflict-free after push, got %v", conflicts) + } +} + +// A duplicate WebSocket read starts from the tracked base but does not hold +// the Syncer's mutex during network I/O. If the watcher pushes a local edit +// before that read is applied, the stale response must not roll the working +// copy or its tracked state back to the old base. +func TestAssessDuplicateWebSocketReadDoesNotEraseCompletedLocalWrite(t *testing.T) { + t.Parallel() + store := newAssessStore(t) + workspaceID := "ws_assess_duplicate_ws_completed_write" + handler := newMountsyncAPIHandler(t, store) + api := httptest.NewServer(handler) + defer api.Close() + + token := mustMountsyncTestJWT(t, "dev-secret", workspaceID, "SandboxB", []string{"fs:read", "fs:write"}, time.Now().Add(time.Hour)) + const remotePath = "/notion/Shared.md" + writeMountsyncRemoteFile(t, api.Client(), api.URL, token, workspaceID, remotePath, "0", "base content") + + localDir := t.TempDir() + wsDisabled := false + syncer, err := NewSyncer(NewHTTPClient(api.URL, token, api.Client()), SyncerOptions{ + WorkspaceID: workspaceID, RemoteRoot: "/notion", LocalRoot: localDir, WebSocket: &wsDisabled, + }) + if err != nil { + t.Fatalf("new syncer: %v", err) + } + ctx := context.Background() + if err := syncer.SyncOnce(ctx); err != nil { + t.Fatalf("bootstrap: %v", err) + } + + blockedClient := &blockingReadFileClient{ + RemoteClient: syncer.client, + started: make(chan struct{}), + release: make(chan struct{}), + } + var releaseOnce sync.Once + releaseBlockedRead := func() { + releaseOnce.Do(func() { close(blockedClient.release) }) + } + defer releaseBlockedRead() + syncer.client = blockedClient + eventResult := make(chan error, 1) + go func() { + eventResult <- syncer.applyWebSocketEvent(ctx, websocketEvent{ + Type: "file.updated", + Path: remotePath, + Timestamp: time.Now().UTC().Format(time.RFC3339Nano), + }) + }() + + select { + case <-blockedClient.started: + case <-time.After(5 * time.Second): + t.Fatal("websocket ReadFile did not reach the interleaving barrier") + } + localPath := filepath.Join(localDir, "Shared.md") + if err := os.WriteFile(localPath, []byte("completed local edit"), 0o644); err != nil { + t.Fatalf("local write: %v", err) + } + if err := syncer.HandleLocalChange(ctx, "Shared.md", 0); err != nil { + t.Fatalf("local watcher push while websocket read was blocked: %v", err) + } + releaseBlockedRead() + select { + case err := <-eventResult: + if err != nil { + t.Fatalf("duplicate websocket apply: %v", err) + } + case <-time.After(5 * time.Second): + t.Fatal("duplicate websocket apply did not finish") + } + + assertLocalFileContent(t, localPath, "completed local edit") + tracked := syncer.state.Files[remotePath] + if tracked.Hash != hashBytes([]byte("completed local edit")) || tracked.Dirty { + t.Fatalf("tracked state rolled back after stale websocket read: %+v", tracked) + } + remoteFile, err := syncer.client.ReadFile(ctx, workspaceID, remotePath) + if err != nil { + t.Fatalf("read server file after interleaving: %v", err) + } + if got := remoteFile.Content; got != "completed local edit" { + t.Fatalf("server content after interleaving = %q, want %q", got, "completed local edit") + } + if conflicts := listConflictArtifacts(t, localDir); len(conflicts) != 0 { + t.Fatalf("duplicate remote base should remain conflict-free, got %v", conflicts) + } +} + // TestAssessSameFileNearSimultaneousWriteGoMergeAutoMerges is the Go-file // counterpart to TestAssessSameFileNearSimultaneousWrite. Both sandboxes race // from one confirmed-clean revision, but their edits target separate diff --git a/internal/mountsync/mount_root_clobber_test.go b/internal/mountsync/mount_root_clobber_test.go index 1d3e8140..01c8cdeb 100644 --- a/internal/mountsync/mount_root_clobber_test.go +++ b/internal/mountsync/mount_root_clobber_test.go @@ -478,9 +478,16 @@ func TestOversizedTrackedFileDoesNotBecomeRemoteDelete(t *testing.T) { if _, ok := fc.files[remotePath]; !ok { t.Fatalf("remote oversized file was deleted") } + info, err := os.Stat(huge) + if err != nil { + t.Fatalf("stat preserved oversized local file: %v", err) + } + if got, want := info.Size(), int64(4096); got != want { + t.Fatalf("oversized local edit was not preserved: size=%d, want %d", got, want) + } state := readPublicState(t, localDir) - if state.PendingWriteback != 0 || state.Files[remotePath].Status != "ready" { - t.Fatalf("oversized tracked drift should settle without writeback: %+v", state) + if state.PendingWriteback != 0 || state.Files[remotePath].Status != "writeback-skipped" || !state.Files[remotePath].Dirty { + t.Fatalf("oversized tracked drift should remain preserved and surface the writeback cap: %+v", state) } } diff --git a/internal/mountsync/syncer.go b/internal/mountsync/syncer.go index 2f0ea684..aa25d930 100644 --- a/internal/mountsync/syncer.go +++ b/internal/mountsync/syncer.go @@ -3320,6 +3320,14 @@ func (s *Syncer) applyWebSocketEvent(ctx context.Context, event websocketEvent) if remotePath == "/" || !isUnderRemoteRoot(s.remoteRoot, remotePath) { return nil } + // ReadFile intentionally runs without mu so local writeback is not + // blocked by remote I/O. Remember the path state before that read: a + // watcher callback can otherwise push a local edit while ReadFile is in + // flight, and the stale response would then use the post-write hash as + // its divergence base and overwrite the newer working copy. + s.mu.Lock() + observedTracked, observedTrackedExists := s.state.Files[remotePath] + s.mu.Unlock() file, err := s.client.ReadFile(ctx, s.workspace, remotePath) if err != nil { var httpErr *HTTPError @@ -3338,6 +3346,13 @@ func (s *Syncer) applyWebSocketEvent(ctx context.Context, event websocketEvent) s.mu.Lock() defer s.mu.Unlock() s.state.LastEventAt = eventAt + currentTracked, currentTrackedExists := s.state.Files[remotePath] + if currentTrackedExists != observedTrackedExists || + (currentTrackedExists && currentTracked != observedTracked) { + s.logf("discarding stale websocket read for %s after tracked state advanced", remotePath) + s.markSyncSuccess() + return s.saveState() + } if err := s.applyRemoteFile(remotePath, file, nil); err != nil { return err } @@ -5999,7 +6014,7 @@ func (s *Syncer) applyRemoteFile(remotePath string, file RemoteFile, conflicted if err != nil { return err } - tracked := s.state.Files[remotePath] + tracked, trackedExists := s.state.Files[remotePath] canWrite := s.canWritePath(remotePath) tracked.ReadOnly = !canWrite tracked.Denied = false @@ -6046,6 +6061,44 @@ func (s *Syncer) applyRemoteFile(remotePath string, file RemoteFile, conflicted localHash := hashBytes(current) if localHash == remoteHash { shouldWrite = false + } else { + // A local filesystem event is debounced for 100ms before + // HandleLocalChange marks the tracked entry Dirty. A competing + // writer's WebSocket update can arrive inside that window. Treat + // on-disk divergence from the last confirmed base as local work + // before replacing the working copy, otherwise the remote update + // silently erases the edit before it ever reaches the outbox CAS. + // + // Revisions are opaque identifiers, not ordered clocks, so decide + // from content alone: both local and remote must diverge from the + // last confirmed base, and from each other. A duplicate remote event + // whose content still equals the base therefore stays harmless. If + // the base hash is unavailable, fail safe and preserve. For an + // untracked path, the simultaneous-create case has no base and uses + // create-only revision "0" in the artifact name. + baseHashUnavailable := trackedExists && tracked.Hash == "" + localDivergedFromBase := !trackedExists || baseHashUnavailable || localHash != tracked.Hash + remoteDivergedFromBase := !trackedExists || baseHashUnavailable || remoteHash != tracked.Hash + if canWrite && !tracked.Dirty && localDivergedFromBase && !remoteDivergedFromBase { + // A duplicate or stale remote event still represents the tracked + // base, so it must not replace a newer local edit or create a + // conflict artifact. The local writeback path remains responsible + // for the local bytes, including its size and policy limits. + shouldWrite = false + } else if canWrite && !tracked.Dirty && localDivergedFromBase && remoteDivergedFromBase { + baseRevision := tracked.Revision + if !trackedExists { + baseRevision = "0" + } + artifactPath, artifactErr := s.writeConflictArtifact(remotePath, baseRevision, current) + if artifactErr != nil { + return artifactErr + } + if conflicted != nil { + conflicted[remotePath] = struct{}{} + } + s.logf("conflict at %s before remote apply; debounced local edit saved at %s", remotePath, artifactPath) + } } } if shouldWrite {