Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 4 additions & 2 deletions internal/mcp/batch_edit_hetero_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -54,19 +54,21 @@ func TestBatchEditItemKind(t *testing.T) {
require.Equal(t, "edit_file", batchEditItem{Path: "p"}.kind(), "a path infers edit_file")
require.Equal(t, "edit_file", batchEditItem{Op: "edit_file", Path: "p"}.kind())
require.Equal(t, "edit_symbol", batchEditItem{Op: "edit_symbol", Path: "p"}.kind(), "explicit op wins over inference")
require.Equal(t, "move_file", batchEditItem{Op: "move_file"}.kind())
require.Equal(t, "delete_file", batchEditItem{Op: "delete_file"}.kind())
}

func TestBatchEditItemsSchemaOneOf(t *testing.T) {
schema := batchEditItemsSchema()
branches, ok := schema["oneOf"].([]any)
require.True(t, ok, "items schema must be a oneOf")
require.Len(t, branches, 2)
require.Len(t, branches, 4)
for _, b := range branches {
m := b.(map[string]any)
require.Equal(t, "object", m["type"])
props := m["properties"].(map[string]any)
op := props["op"].(map[string]any)
require.Contains(t, []any{"edit_symbol", "edit_file"}, op["const"], "each branch is discriminated by an op const")
require.Contains(t, []any{"edit_symbol", "edit_file", "move_file", "delete_file"}, op["const"], "each branch is discriminated by an op const")
require.NotEmpty(t, m["required"], "each branch declares required fields")
}
}
Expand Down
58 changes: 58 additions & 0 deletions internal/mcp/batch_file_lifecycle_recovery_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,58 @@
package mcp

import (
"context"
"errors"
"os"
"path/filepath"
"testing"

"github.com/zzet/gortex/internal/agents"
)

func TestAtomicBatchLifecycleRecoveryRollsBackPartialMove(t *testing.T) {
s := newAtomicBatchTestServer(t, mutationTestWatcher{})
dir := t.TempDir()
source := writeAtomicBatchFixture(t, dir, "source.txt", "source\n")
destination := filepath.Join(dir, "destination.txt")
buffers := map[string]*batchFileBuffer{
source: {
absPath: source, relPath: "source.txt", mode: 0o644,
original: []byte("source\n"), content: []byte("source\n"),
existsBefore: true, existsAfter: false, existenceSet: true,
},
destination: {
absPath: destination, relPath: "destination.txt", mode: 0o644,
content: []byte("source\n"), existsAfter: true, existenceSet: true,
},
}
results := []batchEditResult{{
Op: "move_file", FilePath: "source.txt", DestinationPath: "destination.txt", Status: "validated",
}}
receipt := batchTransactionReceipt{
Version: batchTransactionVersion, TransactionID: "recover-partial-move", Fingerprint: "recovery-fixture",
Status: "preparing", DiskStatus: "unchanged", GraphStatus: "not_started",
Results: results, Summary: batchSummary(results),
}
if err := s.prepareBatchJournal(&receipt, buffers, []string{destination, source}); err != nil {
t.Fatal(err)
}
if err := agents.AtomicWriteFile(destination, []byte("source\n"), 0o644); err != nil {
t.Fatal(err)
}

restarted := &Server{watcher: mutationTestWatcher{}, session: newSessionState()}
recovered, err := restarted.batchTransactionStatus(context.Background(), "recover-partial-move")
if err != nil {
t.Fatal(err)
}
if !recovered.Recovered || recovered.Status != "aborted" || recovered.DiskStatus != "rolled_back" {
t.Fatalf("recovered receipt = %+v", recovered)
}
if got := readAtomicBatchFixture(t, source); got != "source\n" {
t.Fatalf("source = %q", got)
}
if _, err := os.Stat(destination); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("destination survived rollback: %v", err)
}
}
73 changes: 73 additions & 0 deletions internal/mcp/batch_file_lifecycle_security_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,73 @@
package mcp

import (
"context"
"errors"
"os"
"path/filepath"
"strings"
"testing"
)

func TestAtomicBatchLifecycleOutsideRootRefused(t *testing.T) {
repoRoot := t.TempDir()
outsideRoot := t.TempDir()
t.Setenv(batchTransactionDirEnv, filepath.Join(t.TempDir(), "transactions"))
s := newReadGuardServer(t, repoRoot)
source := writeAtomicBatchFixture(t, repoRoot, "source.txt", "source\n")
destination := filepath.Join(outsideRoot, "destination.txt")

receipt, err := s.runBatchTransaction(context.Background(), []batchEditItem{
atomicFileMove(source, destination, ""),
}, "file-lifecycle-outside-root")
if err != nil {
t.Fatal(err)
}
if receipt.Status != "aborted" || receipt.DiskStatus != "unchanged" || !strings.Contains(receipt.Error, "outside") {
t.Fatalf("receipt = %+v", receipt)
}
if got := readAtomicBatchFixture(t, source); got != "source\n" {
t.Fatalf("source = %q", got)
}
if _, err := os.Stat(destination); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("destination was created: %v", err)
}
}

func TestAtomicBatchLifecycleInvalidDigestAndOverlapRefused(t *testing.T) {
t.Run("invalid digest", func(t *testing.T) {
s := newAtomicBatchTestServer(t, mutationTestWatcher{})
path := writeAtomicBatchFixture(t, t.TempDir(), "source.txt", "source\n")
receipt, err := s.runBatchTransaction(context.Background(), []batchEditItem{
atomicFileDelete(path, "not-a-sha256"),
}, "file-lifecycle-invalid-digest")
if err != nil {
t.Fatal(err)
}
if receipt.Status != "aborted" || receipt.DiskStatus != "unchanged" {
t.Fatalf("receipt = %+v", receipt)
}
if got := readAtomicBatchFixture(t, path); got != "source\n" {
t.Fatalf("source = %q", got)
}
})

t.Run("overlapping path ownership", func(t *testing.T) {
s := newAtomicBatchTestServer(t, mutationTestWatcher{})
dir := t.TempDir()
path := writeAtomicBatchFixture(t, dir, "source.txt", "source\n")
receipt, err := s.runBatchTransaction(context.Background(), []batchEditItem{
atomicFileEdit(path, "source", "edited"),
atomicFileDelete(path, ""),
}, "file-lifecycle-overlap")
if err != nil {
t.Fatal(err)
}
if receipt.Status != "aborted" || receipt.DiskStatus != "unchanged" || !strings.Contains(receipt.Error, "overlaps") {
t.Fatalf("receipt = %+v", receipt)
}
if got := readAtomicBatchFixture(t, path); got != "source\n" {
t.Fatalf("source = %q", got)
}
})
}
228 changes: 228 additions & 0 deletions internal/mcp/batch_file_lifecycle_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,228 @@
package mcp

import (
"context"
"crypto/sha256"
"encoding/hex"
"errors"
"os"
"path/filepath"
"strings"
"sync/atomic"
"testing"

"github.com/zzet/gortex/internal/agents"
)

func atomicFileMove(source, destination, expectedSHA256 string) batchEditItem {
return batchEditItem{
Op: "move_file",
SourcePath: source,
DestinationPath: destination,
ExpectedSHA256: expectedSHA256,
}
}

func atomicFileDelete(path, expectedSHA256 string) batchEditItem {
return batchEditItem{Op: "delete_file", Path: path, ExpectedSHA256: expectedSHA256}
}

func testSHA256(content string) string {
sum := sha256.Sum256([]byte(content))
return hex.EncodeToString(sum[:])
}

func TestAtomicBatchMovesAndDeletesWholeFiles(t *testing.T) {
var scheduled atomic.Int64
s := newAtomicBatchTestServer(t, mutationTestWatcher{scheduled: &scheduled})
dir := t.TempDir()
source := writeAtomicBatchFixture(t, dir, "source.txt", "move me\n")
destination := filepath.Join(dir, "nested", "destination.txt")
deleted := writeAtomicBatchFixture(t, dir, "delete.txt", "delete me\n")

receipt, err := s.runBatchTransaction(context.Background(), []batchEditItem{
atomicFileMove(source, destination, testSHA256("move me\n")),
atomicFileDelete(deleted, testSHA256("delete me\n")),
}, "file-lifecycle-success")
if err != nil {
t.Fatal(err)
}
if receipt.Status != "committed" || receipt.DiskStatus != "committed" || receipt.GraphStatus != "fresh" {
t.Fatalf("receipt = %+v", receipt)
}
if _, err := os.Stat(source); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("source still exists or stat failed: %v", err)
}
if got := readAtomicBatchFixture(t, destination); got != "move me\n" {
t.Fatalf("destination = %q", got)
}
if _, err := os.Stat(deleted); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("deleted file still exists or stat failed: %v", err)
}
if scheduled.Load() != 3 {
t.Fatalf("scheduled = %d, want source + destination + deleted", scheduled.Load())
}

states := make(map[string]batchTransactionFile, len(receipt.Files))
for _, file := range receipt.Files {
states[file.Path] = file
}
if !states[source].AfterAbsent || states[source].BeforeAbsent {
t.Fatalf("source state = %+v", states[source])
}
if !states[destination].BeforeAbsent || states[destination].AfterAbsent {
t.Fatalf("destination state = %+v", states[destination])
}
if !states[deleted].AfterAbsent || states[deleted].BeforeAbsent {
t.Fatalf("deleted state = %+v", states[deleted])
}
}

func TestAtomicBatchLifecyclePreconditionsWriteNothing(t *testing.T) {
t.Run("digest mismatch", func(t *testing.T) {
s := newAtomicBatchTestServer(t, mutationTestWatcher{})
dir := t.TempDir()
source := writeAtomicBatchFixture(t, dir, "source.txt", "original\n")
destination := filepath.Join(dir, "destination.txt")

receipt, err := s.runBatchTransaction(context.Background(), []batchEditItem{
atomicFileMove(source, destination, strings.Repeat("0", 64)),
}, "file-lifecycle-digest")
if err != nil {
t.Fatal(err)
}
if receipt.Status != "aborted" || receipt.DiskStatus != "unchanged" {
t.Fatalf("receipt = %+v", receipt)
}
if got := readAtomicBatchFixture(t, source); got != "original\n" {
t.Fatalf("source = %q", got)
}
if _, err := os.Stat(destination); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("destination was created: %v", err)
}
})

t.Run("destination exists", func(t *testing.T) {
s := newAtomicBatchTestServer(t, mutationTestWatcher{})
dir := t.TempDir()
source := writeAtomicBatchFixture(t, dir, "source.txt", "source\n")
destination := writeAtomicBatchFixture(t, dir, "destination.txt", "destination\n")

receipt, err := s.runBatchTransaction(context.Background(), []batchEditItem{
atomicFileMove(source, destination, ""),
}, "file-lifecycle-destination-exists")
if err != nil {
t.Fatal(err)
}
if receipt.Status != "aborted" || receipt.DiskStatus != "unchanged" {
t.Fatalf("receipt = %+v", receipt)
}
if readAtomicBatchFixture(t, source) != "source\n" || readAtomicBatchFixture(t, destination) != "destination\n" {
t.Fatal("precondition failure changed disk")
}
})
}

func TestAtomicBatchLifecycleRollbackRestoresSources(t *testing.T) {
s := newAtomicBatchTestServer(t, mutationTestWatcher{})
dir := t.TempDir()
moveSource := writeAtomicBatchFixture(t, dir, "move-source.txt", "move\n")
moveDestination := filepath.Join(dir, "move-destination.txt")
deleted := writeAtomicBatchFixture(t, dir, "delete.txt", "delete\n")

var removes atomic.Int64
s.batchRemoveOverride = func(path string) error {
if removes.Add(1) == 2 {
return errors.New("injected remove failure")
}
return os.Remove(path)
}

receipt, err := s.runBatchTransaction(context.Background(), []batchEditItem{
atomicFileMove(moveSource, moveDestination, ""),
atomicFileDelete(deleted, ""),
}, "file-lifecycle-rollback")
if err != nil {
t.Fatal(err)
}
if receipt.Status != "aborted" || receipt.DiskStatus != "rolled_back" {
t.Fatalf("receipt = %+v", receipt)
}
if readAtomicBatchFixture(t, moveSource) != "move\n" || readAtomicBatchFixture(t, deleted) != "delete\n" {
t.Fatal("rollback did not restore source files")
}
if _, err := os.Stat(moveDestination); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("rollback did not remove destination: %v", err)
}
}

func TestAtomicBatchLifecycleRejectsSymlinkSource(t *testing.T) {
if os.PathSeparator == '\\' {
t.Skip("symlink creation is not reliably available on Windows CI")
}
s := newAtomicBatchTestServer(t, mutationTestWatcher{})
dir := t.TempDir()
target := writeAtomicBatchFixture(t, dir, "target.txt", "target\n")
link := filepath.Join(dir, "link.txt")
if err := os.Symlink(target, link); err != nil {
t.Fatal(err)
}

receipt, err := s.runBatchTransaction(context.Background(), []batchEditItem{
atomicFileDelete(link, ""),
}, "file-lifecycle-symlink")
if err != nil {
t.Fatal(err)
}
if receipt.Status != "aborted" || !strings.Contains(receipt.Results[0].Error, "symlink") {
t.Fatalf("receipt = %+v", receipt)
}
if got := readAtomicBatchFixture(t, target); got != "target\n" {
t.Fatalf("target = %q", got)
}
}

func TestAtomicBatchLifecycleIdempotentRetry(t *testing.T) {
s := newAtomicBatchTestServer(t, mutationTestWatcher{})
dir := t.TempDir()
source := writeAtomicBatchFixture(t, dir, "source.txt", "source\n")
destination := filepath.Join(dir, "destination.txt")
edits := []batchEditItem{atomicFileMove(source, destination, "")}

first, err := s.runBatchTransaction(context.Background(), edits, "file-lifecycle-idempotent")
if err != nil {
t.Fatal(err)
}
second, err := s.runBatchTransaction(context.Background(), edits, "file-lifecycle-idempotent")
if err != nil {
t.Fatal(err)
}
if first.Status != "committed" || second.Status != "committed" || second.Fingerprint != first.Fingerprint {
t.Fatalf("first=%+v second=%+v", first, second)
}
if got := readAtomicBatchFixture(t, destination); got != "source\n" {
t.Fatalf("destination = %q", got)
}
}

func TestAtomicBatchLifecycleUsesDurableWriterForCreatedDestination(t *testing.T) {
s := newAtomicBatchTestServer(t, mutationTestWatcher{})
dir := t.TempDir()
source := writeAtomicBatchFixture(t, dir, "source.txt", "source\n")
destination := filepath.Join(dir, "destination.txt")
var writes atomic.Int64
s.batchWriteOverride = func(path string, content []byte, mode os.FileMode) error {
writes.Add(1)
return agents.AtomicWriteFile(path, content, mode)
}

receipt, err := s.runBatchTransaction(context.Background(), []batchEditItem{
atomicFileMove(source, destination, ""),
}, "file-lifecycle-writer")
if err != nil {
t.Fatal(err)
}
if receipt.Status != "committed" || writes.Load() != 1 {
t.Fatalf("receipt=%+v writes=%d", receipt, writes.Load())
}
}
Loading
Loading