From ad80180381adb9a18bb3e88555de339c7ebfc396 Mon Sep 17 00:00:00 2001 From: mptyl Date: Tue, 11 Aug 2026 03:13:01 +0200 Subject: [PATCH] fix: harden task1 workspace output and publication --- tools/thothctl/cmd/thothctl/main.go | 94 ++++++++++-- tools/thothctl/cmd/thothctl/main_test.go | 48 +++++- tools/thothctl/internal/safeio/files.go | 3 + tools/thothctl/internal/safeio/files_test.go | 24 +++ tools/thothctl/internal/safeio/files_unix.go | 108 ++++++++++--- .../thothctl/internal/safeio/files_windows.go | 144 ++++++++++++++---- .../internal/safeio/files_windows_test.go | 8 + .../internal/workspaceops/operations.go | 12 +- .../internal/workspaceops/operations_test.go | 36 +++++ 9 files changed, 410 insertions(+), 67 deletions(-) diff --git a/tools/thothctl/cmd/thothctl/main.go b/tools/thothctl/cmd/thothctl/main.go index 2af4fdb7..4fdaf8d4 100644 --- a/tools/thothctl/cmd/thothctl/main.go +++ b/tools/thothctl/cmd/thothctl/main.go @@ -14,6 +14,7 @@ import ( "path/filepath" "strconv" "strings" + "unicode" "github.com/aritmolab/thothii/tools/thothctl/internal/compose" "github.com/aritmolab/thothii/tools/thothctl/internal/config" @@ -68,10 +69,13 @@ func main() { } func run(ctx context.Context, args []string, stdout, stderr io.Writer) int { - // Apply the public stream bounds after all sanitization and final encoding. - // In particular, JSON escaping can expand an otherwise accepted child result. - stdout = &boundedWriter{dst: stdout, maximum: maxPublicStdoutBytes} - stderr = &boundedWriter{dst: stderr, maximum: maxPublicStderrBytes} + // Workspace results are untrusted child output, so only that dispatch receives + // the public bounds. Legacy renderers intentionally retain their established + // behavior and must not silently truncate successful output. + if len(args) > 2 && args[2] == "workspace" { + stdout = &boundedWriter{dst: stdout, maximum: maxPublicStdoutBytes} + stderr = &boundedWriter{dst: stderr, maximum: maxPublicStderrBytes} + } if len(args) == 1 && (args[0] == "--help" || args[0] == "-h") { fmt.Fprint(stdout, usage) return 0 @@ -128,13 +132,19 @@ func run(ctx context.Context, args []string, stdout, stderr io.Writer) int { case workspaceops.RunRequest: jsonMode = c.JSON } + publicResult, projectionErr := projectWorkspaceResult(result, secretValues) + if projectionErr != nil { + fmt.Fprintln(stderr, "thothctl: invalid workspace result") + return 1 + } if jsonMode { - if err := writeWorkspaceJSON(stdout, result); err != nil { + if err := writeWorkspaceJSON(stdout, publicResult); err != nil { fmt.Fprintln(stderr, "thothctl: workspace result exceeds output limit") return 1 } - } else { - renderWorkspaceHuman(stdout, result) + } else if err := renderWorkspaceHuman(stdout, publicResult); err != nil { + fmt.Fprintln(stderr, "thothctl: workspace result exceeds output limit") + return 1 } if result.Status == "blocked" { return 3 @@ -271,18 +281,76 @@ func writeWorkspaceJSON(w io.Writer, result workspaceops.Result) error { return err } -func renderWorkspaceHuman(w io.Writer, result workspaceops.Result) { +func projectWorkspaceResult(result workspaceops.Result, secretValues []string) (workspaceops.Result, error) { + project := func(value string) (string, error) { + value = output.Sanitize(value, secretValues) + for _, r := range value { + if unicode.IsControl(r) { + return "", errors.New("control character in workspace result") + } + } + return value, nil + } + var err error + for _, value := range []*string{&result.Status, &result.Code, &result.WorkspaceID, &result.WorkspaceRevision, &result.DescriptorBlob, &result.Operation, &result.RunID} { + if *value, err = project(*value); err != nil { + return workspaceops.Result{}, err + } + } + if result.ChildRuns != nil { + childRuns := make(map[string]string, len(result.ChildRuns)) + for key, value := range result.ChildRuns { + publicKey, keyErr := project(key) + if keyErr != nil { + return workspaceops.Result{}, keyErr + } + publicValue, valueErr := project(value) + if valueErr != nil { + return workspaceops.Result{}, valueErr + } + childRuns[publicKey] = publicValue + } + result.ChildRuns = childRuns + } + for i := range result.CompletedStages { + if result.CompletedStages[i], err = project(result.CompletedStages[i]); err != nil { + return workspaceops.Result{}, err + } + } + for i := range result.Warnings { + if result.Warnings[i], err = project(result.Warnings[i]); err != nil { + return workspaceops.Result{}, err + } + } + for i := range result.ArtifactIdentities { + if result.ArtifactIdentities[i].Kind, err = project(result.ArtifactIdentities[i].Kind); err != nil { + return workspaceops.Result{}, err + } + if result.ArtifactIdentities[i].Digest, err = project(result.ArtifactIdentities[i].Digest); err != nil { + return workspaceops.Result{}, err + } + } + return result, nil +} + +func renderWorkspaceHuman(w io.Writer, result workspaceops.Result) error { if result.Code == workspaceops.CodeRegistryBootstrapRecoveryConflict { - fmt.Fprintln(w, "Bootstrap recovery is ambiguous or corrupt; inspect the installation registry jobs.") - return + _, err := fmt.Fprintln(w, "Bootstrap recovery is ambiguous or corrupt; inspect the installation registry jobs.") + return err + } + if _, err := fmt.Fprintf(w, "Workspace %s: %s (%s)\n", result.WorkspaceID, result.Status, result.Code); err != nil { + return err } - fmt.Fprintf(w, "Workspace %s: %s (%s)\n", result.WorkspaceID, result.Status, result.Code) if result.RunID != "" { - fmt.Fprintf(w, "Run: %s\n", result.RunID) + if _, err := fmt.Fprintf(w, "Run: %s\n", result.RunID); err != nil { + return err + } } if len(result.CompletedStages) > 0 { - fmt.Fprintf(w, "Completed stages: %s\n", strings.Join(result.CompletedStages, ", ")) + _, err := fmt.Fprintf(w, "Completed stages: %s\n", strings.Join(result.CompletedStages, ", ")) + return err } + return nil } func writeRemovalTargets(outputWriter io.Writer, project string, targets []serverops.Container) { diff --git a/tools/thothctl/cmd/thothctl/main_test.go b/tools/thothctl/cmd/thothctl/main_test.go index 4c255ee5..bab97c9a 100644 --- a/tools/thothctl/cmd/thothctl/main_test.go +++ b/tools/thothctl/cmd/thothctl/main_test.go @@ -878,6 +878,50 @@ func TestRunPiStatusPreservesDockerExitCodeAndRedactsDiagnostics(t *testing.T) { } } +func TestRunWorkspacePublicProjectionRedactsSecrets(t *testing.T) { + fixture := newCLIFixture(t, "TOKEN_FILE=%s\n") + secretPath := filepath.Join(fixture.root, "token") + if err := os.WriteFile(secretPath, []byte("workspace-secret"), 0o600); err != nil { + t.Fatal(err) + } + fixture.setEnvironment(t, secretPath) + t.Setenv("THOTHCTL_FAKE_WORKSPACE_RESULT", fmt.Sprintf(`{"schemaVersion":1,"status":"succeeded","code":"ok","workspaceId":"psd","workspaceRevision":"%s","descriptorBlob":"%s","operation":"inspect","completedStages":["stage workspace-secret"],"warnings":["warning workspace-secret"]}`, strings.Repeat("0", 40), strings.Repeat("a", 40))) + var stdout, stderr bytes.Buffer + if code := run(context.Background(), []string{"--installation", fixture.installationPath, "workspace", "inspect", "--workspace", "psd", "--json"}, &stdout, &stderr); code != 0 { + t.Fatalf("run() exit=%d stderr=%q", code, stderr.String()) + } + if strings.Contains(stdout.String(), "workspace-secret") || !strings.Contains(stdout.String(), "[REDACTED]") { + t.Fatalf("public JSON = %q, want redacted secret", stdout.String()) + } +} + +func TestRunWorkspacePublicProjectionRejectsHumanControlCharacters(t *testing.T) { + fixture := newCLIFixture(t, "") + fixture.setEnvironment(t) + t.Setenv("THOTHCTL_FAKE_WORKSPACE_RESULT", fmt.Sprintf(`{"schemaVersion":1,"status":"succeeded","code":"ok","workspaceId":"psd","workspaceRevision":"%s","descriptorBlob":"%s","operation":"inspect","completedStages":["safe\nforged"]}`, strings.Repeat("0", 40), strings.Repeat("a", 40))) + var stdout, stderr bytes.Buffer + if code := run(context.Background(), []string{"--installation", fixture.installationPath, "workspace", "inspect", "--workspace", "psd"}, &stdout, &stderr); code != 1 || stdout.Len() != 0 { + t.Fatalf("run() exit=%d stdout=%q stderr=%q, want rejected output", code, stdout.String(), stderr.String()) + } +} + +func TestRunNonWorkspaceOutputIsNotGloballyCapped(t *testing.T) { + fixture := newCLIFixture(t, "") + fixture.setEnvironment(t) + logPath := filepath.Join(fixture.root, "large.log") + if err := os.WriteFile(logPath, []byte(strings.Repeat("log-line\n", 200000)), 0o600); err != nil { + t.Fatal(err) + } + t.Setenv("THOTHCTL_FAKE_LOG_FILE", logPath) + var stdout, stderr bytes.Buffer + if code := run(context.Background(), []string{"--installation", fixture.installationPath, "logs"}, &stdout, &stderr); code != 0 { + t.Fatalf("run() exit=%d stderr=%q invocations=%#v", code, stderr.String(), fixture.invocations(t)) + } + if stdout.Len() <= maxPublicStdoutBytes { + t.Fatalf("non-workspace output length=%d, want greater than cap %d", stdout.Len(), maxPublicStdoutBytes) + } +} + type cliFixture struct { root string installationPath string @@ -951,7 +995,8 @@ case " $* " in *"settings-cli.js --snapshot"*) printf '%s\n' '{"exists":false,"rawBase64":""}' ;; *"/settings "*) printf '%s\n' '{"provider":"provider","model":"model","thinking":"medium"}' ;; *"/internal/maintenance/status "*) printf '%s\n' '{"active":true,"admissions":0}' ;; - *" logs "*) printf '%s\n' "$THOTHCTL_FAKE_LOG" ;; + *" logs "*) + if [ -n "${THOTHCTL_FAKE_LOG_FILE:-}" ]; then /bin/cat "$THOTHCTL_FAKE_LOG_FILE"; else printf '%s\n' "$THOTHCTL_FAKE_LOG"; fi ;; esac if [ "${THOTHCTL_FAKE_FAIL_ON:-}" = "version" ]; then printf '%s\n' "${THOTHCTL_FAKE_FAILURE:-fake Docker failure}" >&2 @@ -986,6 +1031,7 @@ func (f cliFixture) setEnvContents(t *testing.T, env string) { t.Setenv("THOTHCTL_FAKE_ARGS", f.argsFile) t.Setenv("THOTHCTL_FAKE_EXIT", "0") t.Setenv("THOTHCTL_FAKE_LOG", "") + t.Setenv("THOTHCTL_FAKE_LOG_FILE", "") t.Setenv("THOTHCTL_FAKE_FAILURE", "") t.Setenv("THOTHCTL_FAKE_FAIL_ON", "") t.Setenv("THOTHCTL_FAKE_STOPPED_PS", "[]") diff --git a/tools/thothctl/internal/safeio/files.go b/tools/thothctl/internal/safeio/files.go index 5d853bc8..ea1c4949 100644 --- a/tools/thothctl/internal/safeio/files.go +++ b/tools/thothctl/internal/safeio/files.go @@ -18,6 +18,9 @@ func ValidateCanonicalPath(path string) error { if !filepath.IsAbs(path) || filepath.Clean(path) != path || strings.Contains(path, string(filepath.Separator)+".."+string(filepath.Separator)) { return ErrUnsafeFile } + if err := validatePlatformPathSyntax(path); err != nil { + return err + } return nil } diff --git a/tools/thothctl/internal/safeio/files_test.go b/tools/thothctl/internal/safeio/files_test.go index a2033628..6c0d798b 100644 --- a/tools/thothctl/internal/safeio/files_test.go +++ b/tools/thothctl/internal/safeio/files_test.go @@ -4,6 +4,7 @@ import ( "errors" "os" "path/filepath" + "strings" "testing" "github.com/aritmolab/thothii/tools/thothctl/internal/testsupport" @@ -86,6 +87,29 @@ func TestWriteCanonicalExclusiveRejectsExistingAndCreatesPrivateFile(t *testing. } } +func TestWriteCanonicalExclusiveCleansPrivateStageWhenPublicationRacesExistingLeaf(t *testing.T) { + root, err := filepath.EvalSymlinks(t.TempDir()) + if err != nil { + t.Fatal(err) + } + path := filepath.Join(root, "candidate.yaml") + if err := os.WriteFile(path, []byte("attacker"), 0o600); err != nil { + t.Fatal(err) + } + if err := WriteCanonicalExclusive(path, []byte("candidate"), 0o600); !errors.Is(err, ErrUnsafeFile) { + t.Fatalf("write error=%v", err) + } + entries, err := os.ReadDir(root) + if err != nil { + t.Fatal(err) + } + for _, entry := range entries { + if strings.HasPrefix(entry.Name(), ".thothctl-candidate-") { + t.Fatalf("private stage leaked: %s", entry.Name()) + } + } +} + func TestReadCanonicalRegularRejectsHardlinkAndDirectory(t *testing.T) { root, err := filepath.EvalSymlinks(t.TempDir()) if err != nil { diff --git a/tools/thothctl/internal/safeio/files_unix.go b/tools/thothctl/internal/safeio/files_unix.go index 1a257a30..2aa487c6 100644 --- a/tools/thothctl/internal/safeio/files_unix.go +++ b/tools/thothctl/internal/safeio/files_unix.go @@ -3,6 +3,8 @@ package safeio import ( + "crypto/rand" + "encoding/hex" "io/fs" "os" "strings" @@ -90,51 +92,117 @@ func writeCanonicalExclusive(path string, contents []byte, mode fs.FileMode) err if err != nil { return ErrUnsafeFile } - defer func() { unix.Close(dir) }() + defer unix.Close(dir) for _, component := range components[:len(components)-1] { - next, err := unix.Openat(dir, component, unix.O_RDONLY|unix.O_DIRECTORY|unix.O_CLOEXEC|unix.O_NOFOLLOW, 0) - if err != nil { + next, openErr := unix.Openat(dir, component, unix.O_RDONLY|unix.O_DIRECTORY|unix.O_CLOEXEC|unix.O_NOFOLLOW, 0) + if openErr != nil { return ErrUnsafeFile } unix.Close(dir) dir = next } - var parentBefore unix.Stat_t - if err := unix.Fstat(dir, &parentBefore); err != nil { + if !recheckUnixParentPath(components[:len(components)-1], dir) { return ErrUnsafeFile } - fd, err := unix.Openat(dir, components[len(components)-1], unix.O_WRONLY|unix.O_CREAT|unix.O_EXCL|unix.O_CLOEXEC|unix.O_NOFOLLOW, uint32(mode.Perm())) + + // Build the candidate under a private, same-parent name. Only after it is fully + // written, synced, and identity-checked do we link it into the requested leaf. + stage, err := privateStageName() if err != nil { return ErrUnsafeFile } - f := os.NewFile(uintptr(fd), "thothctl-safeio-output") - if f == nil { - unix.Close(fd) + stageFD, err := unix.Openat(dir, stage, unix.O_WRONLY|unix.O_CREAT|unix.O_EXCL|unix.O_CLOEXEC|unix.O_NOFOLLOW, uint32(mode.Perm())) + if err != nil { return ErrUnsafeFile } - defer f.Close() - if err := f.Chmod(mode); err != nil { + stageFile := os.NewFile(uintptr(stageFD), "thothctl-safeio-stage") + if stageFile == nil { + unix.Close(stageFD) + _ = unix.Unlinkat(dir, stage, 0) return ErrUnsafeFile } - if _, err := f.Write(contents); err != nil { + stageCreated := true + var staged unix.Stat_t + if err := unix.Fstat(stageFD, &staged); err != nil || staged.Nlink != 1 || staged.Mode&unix.S_IFMT != unix.S_IFREG { + _ = stageFile.Close() + _ = unix.Unlinkat(dir, stage, 0) return ErrUnsafeFile } - if err := f.Sync(); err != nil { + published := false + var publishedIdentity unix.Stat_t + cleanup := func() { + if published { + var current unix.Stat_t + if unix.Fstatat(dir, components[len(components)-1], ¤t, unix.AT_SYMLINK_NOFOLLOW) == nil && + current.Ino == publishedIdentity.Ino && current.Dev == publishedIdentity.Dev { + _ = unix.Unlinkat(dir, components[len(components)-1], 0) + } + } + if stageCreated { + var current unix.Stat_t + if unix.Fstatat(dir, stage, ¤t, unix.AT_SYMLINK_NOFOLLOW) == nil && + current.Ino == staged.Ino && current.Dev == staged.Dev { + _ = unix.Unlinkat(dir, stage, 0) + } + } + } + fail := func() error { + _ = stageFile.Close() + cleanup() return ErrUnsafeFile } - var opened, named, parentAfter unix.Stat_t - if err := unix.Fstat(fd, &opened); err != nil || opened.Nlink != 1 { + if err := stageFile.Chmod(mode); err != nil { + return fail() + } + n, err := stageFile.Write(contents) + if err != nil || n != len(contents) { + return fail() + } + if err := stageFile.Sync(); err != nil { + return fail() + } + var afterWrite unix.Stat_t + if err := unix.Fstat(stageFD, &afterWrite); err != nil || afterWrite.Nlink != 1 || afterWrite.Mode&unix.S_IFMT != unix.S_IFREG || afterWrite.Size != int64(len(contents)) { + return fail() + } + if err := stageFile.Close(); err != nil { + cleanup() return ErrUnsafeFile } - if err := unix.Fstatat(dir, components[len(components)-1], &named, unix.AT_SYMLINK_NOFOLLOW); err != nil || named.Nlink != 1 || named.Ino != opened.Ino || named.Dev != opened.Dev { - return ErrUnsafeFile + stageCreated = true + if !recheckUnixParentPath(components[:len(components)-1], dir) { + return fail() } - if err := unix.Fstat(dir, &parentAfter); err != nil || parentAfter.Ino != parentBefore.Ino || parentAfter.Dev != parentBefore.Dev || !recheckUnixParentPath(components[:len(components)-1], dir) { - return ErrUnsafeFile + if err := unix.Linkat(dir, stage, dir, components[len(components)-1], 0); err != nil { + return fail() + } + published = true + publishedIdentity = staged + if err := unix.Fstatat(dir, components[len(components)-1], &publishedIdentity, unix.AT_SYMLINK_NOFOLLOW); err != nil || + publishedIdentity.Ino != staged.Ino || publishedIdentity.Dev != staged.Dev || publishedIdentity.Nlink != 2 { + return fail() + } + if err := unix.Unlinkat(dir, stage, 0); err != nil { + return fail() + } + stageCreated = false + if err := unix.Fstatat(dir, components[len(components)-1], &publishedIdentity, unix.AT_SYMLINK_NOFOLLOW); err != nil || publishedIdentity.Nlink != 1 { + return fail() + } + if err := unix.Fsync(dir); err != nil || !recheckUnixParentPath(components[:len(components)-1], dir) { + return fail() } return nil } +func privateStageName() (string, error) { + var random [16]byte + if _, err := rand.Read(random[:]); err != nil { + return "", err + } + return ".thothctl-candidate-" + hex.EncodeToString(random[:]), nil +} + func validateCanonicalOutputPath(path string) error { if err := ValidateCanonicalPath(path); err != nil { return err @@ -207,3 +275,5 @@ func recheckUnixParentPath(components []string, retained int) bool { var got, want unix.Stat_t return unix.Fstat(dir, &got) == nil && unix.Fstat(retained, &want) == nil && got.Ino == want.Ino && got.Dev == want.Dev } + +func validatePlatformPathSyntax(path string) error { return nil } diff --git a/tools/thothctl/internal/safeio/files_windows.go b/tools/thothctl/internal/safeio/files_windows.go index 810cb400..00209238 100644 --- a/tools/thothctl/internal/safeio/files_windows.go +++ b/tools/thothctl/internal/safeio/files_windows.go @@ -3,6 +3,8 @@ package safeio import ( + "crypto/rand" + "encoding/hex" "io/fs" "os" "path/filepath" @@ -137,50 +139,113 @@ func writeCanonicalExclusive(path string, contents []byte, mode fs.FileMode) err return ErrUnsafeFile } defer closeWindowsHandles(retainedParents) - // Parent handles stay open with delete sharing denied until publication and - // identity recheck complete; this is the Windows equivalent of retained dirfds. securityDescriptor, securityAttributes, err := ownerOnlySecurityAttributes() if err != nil { return ErrUnsafeFile } - h, err := windows.CreateFile(windows.StringToUTF16Ptr(filepath.Join(parent, filepath.Base(path))), windows.GENERIC_WRITE, windowsOutputHandleShareMode, securityAttributes, windows.CREATE_NEW, windows.FILE_ATTRIBUTE_NORMAL|windows.FILE_FLAG_OPEN_REPARSE_POINT, 0) - if err != nil { - return ErrUnsafeFile - } - f := os.NewFile(uintptr(h), "thothctl-safeio-output") - if f == nil { - windows.CloseHandle(h) - return ErrUnsafeFile - } - defer f.Close() - // Go's Windows Chmod only toggles the read-only attribute; it cannot enforce - // owner-only permissions. The restrictive DACL is installed at creation time - // through SECURITY_ATTRIBUTES above. _ = securityDescriptor - if mode.Perm() != 0o600 { - return ErrUnsafeFile - } - if _, err := f.Write(contents); err != nil { - return ErrUnsafeFile - } - if err := f.Sync(); err != nil { - return ErrUnsafeFile - } - var opened, named windows.ByHandleFileInformation - if windows.GetFileInformationByHandle(h, &opened) != nil || opened.NumberOfLinks != 1 { - return ErrUnsafeFile - } - check, err := openWindowsComponent(path, false) + stageName, err := privateWindowsStageName() if err != nil { return ErrUnsafeFile } - defer windows.CloseHandle(check) - if windows.GetFileInformationByHandle(check, &named) != nil || named.NumberOfLinks != 1 || !sameWindowsFile(opened, named) { + stagePath := filepath.Join(parent, stageName) + stageHandle, err := windows.CreateFile(windows.StringToUTF16Ptr(stagePath), windows.GENERIC_WRITE, windowsOutputHandleShareMode, securityAttributes, windows.CREATE_NEW, windows.FILE_ATTRIBUTE_NORMAL|windows.FILE_FLAG_OPEN_REPARSE_POINT, 0) + if err != nil { return ErrUnsafeFile } + stageFile := os.NewFile(uintptr(stageHandle), "thothctl-safeio-stage") + if stageFile == nil { + windows.CloseHandle(stageHandle) + _ = windows.DeleteFile(windows.StringToUTF16Ptr(stagePath)) + return ErrUnsafeFile + } + stageCreated := true + published := false + var staged, publishedIdentity windows.ByHandleFileInformation + if err := windows.GetFileInformationByHandle(stageHandle, &staged); err != nil || staged.NumberOfLinks != 1 { + _ = stageFile.Close() + _ = windows.DeleteFile(windows.StringToUTF16Ptr(stagePath)) + return ErrUnsafeFile + } + cleanup := func() { + if published { + removeWindowsIfIdentity(filepath.Join(parent, filepath.Base(path)), publishedIdentity) + } + if stageCreated { + removeWindowsIfIdentity(stagePath, staged) + } + } + fail := func() error { + _ = stageFile.Close() + cleanup() + return ErrUnsafeFile + } + if n, err := stageFile.Write(contents); err != nil || n != len(contents) { + return fail() + } + if err := stageFile.Sync(); err != nil { + return fail() + } + var afterWrite windows.ByHandleFileInformation + if err := windows.GetFileInformationByHandle(stageHandle, &afterWrite); err != nil || afterWrite.NumberOfLinks != 1 || afterWrite.FileSizeHigh != uint32(uint64(len(contents))>>32) || afterWrite.FileSizeLow != uint32(len(contents)) { + return fail() + } + if err := stageFile.Close(); err != nil { + cleanup() + return ErrUnsafeFile + } + // CreateHardLink is an atomic, same-volume, no-replace publication. The final + // pathname can never refer to a partially written candidate. + finalPath := filepath.Join(parent, filepath.Base(path)) + if err := windows.CreateHardLink(windows.StringToUTF16Ptr(finalPath), windows.StringToUTF16Ptr(stagePath), 0); err != nil { + return fail() + } + published = true + publishedIdentity = staged + check, identityErr := windowsFileIdentity(finalPath) + if identityErr != nil || check.NumberOfLinks != 2 || !sameWindowsFile(staged, check) { + return fail() + } + publishedIdentity = check + if err := windows.DeleteFile(windows.StringToUTF16Ptr(stagePath)); err != nil { + return fail() + } + stageCreated = false + finalIdentity, identityErr := windowsFileIdentity(finalPath) + if identityErr != nil || finalIdentity.NumberOfLinks != 1 || !sameWindowsFile(staged, finalIdentity) { + return fail() + } return nil } +func privateWindowsStageName() (string, error) { + var random [16]byte + if _, err := rand.Read(random[:]); err != nil { + return "", err + } + return ".thothctl-candidate-" + hex.EncodeToString(random[:]), nil +} + +func windowsFileIdentity(path string) (windows.ByHandleFileInformation, error) { + h, err := windows.CreateFile(windows.StringToUTF16Ptr(path), windows.GENERIC_READ, windows.FILE_SHARE_READ|windows.FILE_SHARE_WRITE, nil, windows.OPEN_EXISTING, windows.FILE_ATTRIBUTE_NORMAL|windows.FILE_FLAG_OPEN_REPARSE_POINT, 0) + if err != nil { + return windows.ByHandleFileInformation{}, err + } + defer windows.CloseHandle(h) + var information windows.ByHandleFileInformation + if err := windows.GetFileInformationByHandle(h, &information); err != nil || information.NumberOfLinks == 0 || information.FileAttributes&windows.FILE_ATTRIBUTE_REPARSE_POINT != 0 || information.FileAttributes&windows.FILE_ATTRIBUTE_DIRECTORY != 0 { + return windows.ByHandleFileInformation{}, ErrUnsafeFile + } + return information, nil +} + +func removeWindowsIfIdentity(path string, expected windows.ByHandleFileInformation) { + got, err := windowsFileIdentity(path) + if err == nil && sameWindowsFile(got, expected) { + _ = windows.DeleteFile(windows.StringToUTF16Ptr(path)) + } +} + func openWindowsParents(path string) (string, []windows.Handle, error) { volume := filepath.VolumeName(path) root := volume + string(filepath.Separator) @@ -217,6 +282,23 @@ func validateCanonicalOutputPath(path string) error { return nil } +func validatePlatformPathSyntax(path string) error { + // Win32 device namespaces and alternate data streams do not represent an + // independent regular file with owner-only output permissions. + lower := strings.ToLower(path) + if strings.HasPrefix(lower, `\\?\`) || strings.HasPrefix(lower, `\\.\`) || + strings.HasPrefix(lower, `\device\`) || strings.HasPrefix(lower, `\??\`) || strings.HasPrefix(path, `\\`) { + return ErrUnsafeFile + } + volume := filepath.VolumeName(path) + for _, component := range strings.Split(strings.TrimPrefix(path, volume+string(filepath.Separator)), string(filepath.Separator)) { + if strings.Contains(component, ":") { + return ErrUnsafeFile + } + } + return nil +} + func sameWindowsFile(a, b windows.ByHandleFileInformation) bool { return a.VolumeSerialNumber == b.VolumeSerialNumber && a.FileIndexHigh == b.FileIndexHigh && a.FileIndexLow == b.FileIndexLow } diff --git a/tools/thothctl/internal/safeio/files_windows_test.go b/tools/thothctl/internal/safeio/files_windows_test.go index 3f3cdfa8..e7457a58 100644 --- a/tools/thothctl/internal/safeio/files_windows_test.go +++ b/tools/thothctl/internal/safeio/files_windows_test.go @@ -22,6 +22,14 @@ var _ [expectedWindowsRetainedHandleShareMode - windowsRetainedHandleShareMode]s var _ [windowsOutputHandleShareMode - expectedWindowsOutputHandleShareMode]struct{} var _ [expectedWindowsOutputHandleShareMode - windowsOutputHandleShareMode]struct{} +func TestValidateCanonicalPathRejectsWindowsNamespacesAndAlternateStreams(t *testing.T) { + for _, path := range []string{`C:\dir\existing.txt:candidate`, `\\?\C:\dir\candidate`, `\\.\pipe\candidate`, `\Device\HarddiskVolume1\candidate`, `\??\C:\candidate`} { + if err := ValidateCanonicalPath(path); !errors.Is(err, ErrUnsafeFile) { + t.Errorf("ValidateCanonicalPath(%q) = %v, want ErrUnsafeFile", path, err) + } + } +} + func TestOpenWindowsComponentBlocksMutationWhileHandleIsRetained(t *testing.T) { t.Run("parent rename", func(t *testing.T) { parent := filepath.Join(t.TempDir(), "parent") diff --git a/tools/thothctl/internal/workspaceops/operations.go b/tools/thothctl/internal/workspaceops/operations.go index 81d2fa23..7d7b2179 100644 --- a/tools/thothctl/internal/workspaceops/operations.go +++ b/tools/thothctl/internal/workspaceops/operations.go @@ -11,6 +11,7 @@ import ( "errors" "fmt" "io" + "path/filepath" "reflect" "regexp" "strings" @@ -369,7 +370,7 @@ func makeInput(c Command) (inputEnvelope, string, error) { if total > maxSQLTotal { return env, "", errors.New("SQL input exceeds limit") } - base := path[strings.LastIndexAny(path, "/\\")+1:] + base := filepath.Base(path) env.SQL = append(env.SQL, sqlInput{base, base64.StdEncoding.EncodeToString(b), DigestBytes(b)}) } case CheckSchemaRequest: @@ -379,7 +380,7 @@ func makeInput(c Command) (inputEnvelope, string, error) { if e != nil { return env, "", errors.New("unsafe annotation input") } - base := x.Annotations[strings.LastIndexAny(x.Annotations, "/\\")+1:] + base := filepath.Base(x.Annotations) env.Annotations = &annotationInput{base, base64.StdEncoding.EncodeToString(b), DigestBytes(b)} } case IndexSchemaRequest: @@ -589,7 +590,12 @@ func publishCandidate(x *hostExport, result Result, path string) error { return errors.New("invalid candidate export") } var doc any - if yaml.Unmarshal(b, &doc) != nil { + decoder := yaml.NewDecoder(bytes.NewReader(b)) + if decoder.Decode(&doc) != nil { + return errors.New("invalid candidate export") + } + var trailing any + if err := decoder.Decode(&trailing); err != io.EOF { return errors.New("invalid candidate export") } if result.RunID == "" || !runIDPattern.MatchString(result.RunID) { diff --git a/tools/thothctl/internal/workspaceops/operations_test.go b/tools/thothctl/internal/workspaceops/operations_test.go index cff3cae7..4c887c0e 100644 --- a/tools/thothctl/internal/workspaceops/operations_test.go +++ b/tools/thothctl/internal/workspaceops/operations_test.go @@ -3,6 +3,7 @@ package workspaceops import ( "bytes" "encoding/base64" + "encoding/json" "fmt" "os" "path/filepath" @@ -16,6 +17,41 @@ import ( "github.com/aritmolab/thothii/tools/thothctl/internal/safeio" ) +func TestMakeInputUsesPlatformBasename(t *testing.T) { + if filepath.Separator == '\\' { + t.Skip("Unix basename behavior") + } + root, err := filepath.EvalSymlinks(t.TempDir()) + if err != nil { + t.Fatal(err) + } + path := filepath.Join(root, "schema\\v2.sql") + if err := os.WriteFile(path, []byte("select 1"), 0o600); err != nil { + t.Fatal(err) + } + _, payload, err := makeInput(SuggestFksRequest{WorkspaceID: "psd", FromSQL: []string{path}}) + if err != nil { + t.Fatal(err) + } + var envelope inputEnvelope + if err := json.Unmarshal([]byte(payload), &envelope); err != nil { + t.Fatal(err) + } + if len(envelope.SQL) != 1 || envelope.SQL[0].Basename != "schema\\v2.sql" { + t.Fatalf("basename=%q", envelope.SQL[0].Basename) + } +} + +func TestPublishCandidateRejectsTrailingYAMLDocument(t *testing.T) { + candidate := []byte("candidates: []\n---\n[\n") + digest := DigestBytes(candidate) + x := hostExport{MediaType: "application/yaml", SHA256: digest, ContentBase64: base64.StdEncoding.EncodeToString(candidate)} + result := Result{RunID: strings.Repeat("b", 32)} + if err := publishCandidate(&x, result, ""); err == nil { + t.Fatal("accepted malformed trailing YAML document") + } +} + func TestParseWorkspaceCommands(t *testing.T) { cases := []struct { name string