diff --git a/reparse.go b/reparse.go index b03e517..9847fbd 100644 --- a/reparse.go +++ b/reparse.go @@ -5,6 +5,7 @@ package winio import ( "bytes" "encoding/binary" + "errors" "fmt" "strings" "unicode/utf16" @@ -14,6 +15,9 @@ import ( const ( reparseTagMountPoint = 0xA0000003 reparseTagSymlink = 0xA000000C + reparseTagLxSymlink = 0xA000001D // WSL/MSYS2 native symlinks + + lxSymlinkVersion = 2 // LX symlink format version ) type reparseDataBuffer struct { @@ -30,6 +34,7 @@ type reparseDataBuffer struct { type ReparsePoint struct { Target string IsMountPoint bool + IsLxSymlink bool // True if this is an LX symlink (WSL/MSYS2 native) } // UnsupportedReparsePointError is returned when trying to decode a non-symlink or @@ -50,14 +55,19 @@ func DecodeReparsePoint(b []byte) (*ReparsePoint, error) { } func DecodeReparsePointData(tag uint32, b []byte) (*ReparsePoint, error) { - isMountPoint := false switch tag { case reparseTagMountPoint: - isMountPoint = true + return decodeWindowsReparsePointData(b, true) case reparseTagSymlink: + return decodeWindowsReparsePointData(b, false) + case reparseTagLxSymlink: + return decodeLxReparsePointData(b) default: return nil, &UnsupportedReparsePointError{tag} } +} + +func decodeWindowsReparsePointData(b []byte, isMountPoint bool) (*ReparsePoint, error) { nameOffset := 8 + binary.LittleEndian.Uint16(b[4:6]) if !isMountPoint { nameOffset += 4 @@ -68,16 +78,56 @@ func DecodeReparsePointData(tag uint32, b []byte) (*ReparsePoint, error) { if err != nil { return nil, err } - return &ReparsePoint{string(utf16.Decode(name)), isMountPoint}, nil + return &ReparsePoint{Target: string(utf16.Decode(name)), IsMountPoint: isMountPoint, IsLxSymlink: false}, nil +} + +func decodeLxReparsePointData(b []byte) (*ReparsePoint, error) { + // LX symlinks store the target as UTF-8 after a 4-byte version field + if len(b) < 4 { + return nil, errors.New("LX symlink buffer too short") + } + targetBytes := b[4:] + for i, c := range targetBytes { + if c == 0 { + targetBytes = targetBytes[:i] + break + } + } + target := string(targetBytes) + return &ReparsePoint{Target: target, IsMountPoint: false, IsLxSymlink: true}, nil } func isDriveLetter(c byte) bool { return (c >= 'a' && c <= 'z') || (c >= 'A' && c <= 'Z') } -// EncodeReparsePoint encodes a Win32 REPARSE_DATA_BUFFER structure describing a symlink or -// mount point. +// EncodeReparsePoint encodes a Win32 REPARSE_DATA_BUFFER structure describing a symlink, +// mount point, or LX symlink. func EncodeReparsePoint(rp *ReparsePoint) []byte { + if rp == nil { + return nil + } + if rp.IsLxSymlink { + return encodeLxReparsePoint(rp) + } + return encodeWindowsReparsePoint(rp) +} + +func encodeLxReparsePoint(rp *ReparsePoint) []byte { + // LX symlink: 4-byte version + UTF-8 target + targetBytes := []byte(rp.Target) + dataLength := 4 + len(targetBytes) + + var b bytes.Buffer + _ = binary.Write(&b, binary.LittleEndian, uint32(reparseTagLxSymlink)) + _ = binary.Write(&b, binary.LittleEndian, uint16(dataLength)) + _ = binary.Write(&b, binary.LittleEndian, uint16(0)) + _ = binary.Write(&b, binary.LittleEndian, uint32(lxSymlinkVersion)) + _, _ = b.Write(targetBytes) + return b.Bytes() +} + +func encodeWindowsReparsePoint(rp *ReparsePoint) []byte { // Generate an NT path and determine if this is a relative path. var ntTarget string relative := false diff --git a/reparse_lx_test.go b/reparse_lx_test.go new file mode 100644 index 0000000..f0c9c4c --- /dev/null +++ b/reparse_lx_test.go @@ -0,0 +1,152 @@ +//go:build windows + +package winio + +import ( + "testing" +) + +const ( + testLxSymlinkAbsolutePath = "/usr/bin/bash" + testWindowsSymlinkPath = `C:\Windows\System32` + testLxSymlinkRelativePath = "../bin/sh" + testLxSymlinkSpecialCharsPath = "/path/with spaces/and-special!@#$%/файл.txt" +) + +func TestLxSymlinkRoundTrip(t *testing.T) { + // Test LX symlink encode/decode + original := &ReparsePoint{ + Target: testLxSymlinkAbsolutePath, + IsMountPoint: false, + IsLxSymlink: true, + } + + // Encode + encoded := EncodeReparsePoint(original) + + // Decode + decoded, err := DecodeReparsePoint(encoded) + if err != nil { + t.Fatalf("Failed to decode: %v", err) + } + + // Verify + if decoded.Target != original.Target { + t.Errorf("Target mismatch: got %q, want %q", decoded.Target, original.Target) + } + if decoded.IsLxSymlink != original.IsLxSymlink { + t.Errorf("IsLxSymlink mismatch: got %v, want %v", decoded.IsLxSymlink, original.IsLxSymlink) + } + if decoded.IsMountPoint != original.IsMountPoint { + t.Errorf("IsMountPoint mismatch: got %v, want %v", decoded.IsMountPoint, original.IsMountPoint) + } +} + +func TestWindowsSymlinkNotLx(t *testing.T) { + // Test that regular Windows symlinks are not marked as LX + original := &ReparsePoint{ + Target: testWindowsSymlinkPath, + IsMountPoint: false, + IsLxSymlink: false, + } + + // Encode + encoded := EncodeReparsePoint(original) + + // Decode + decoded, err := DecodeReparsePoint(encoded) + if err != nil { + t.Fatalf("Failed to decode: %v", err) + } + + // Verify it's NOT an LX symlink + if decoded.IsLxSymlink { + t.Errorf("Windows symlink incorrectly marked as LX symlink") + } +} + +func TestLxSymlinkEmptyTarget(t *testing.T) { + // Test LX symlink with empty target + original := &ReparsePoint{ + Target: "", + IsMountPoint: false, + IsLxSymlink: true, + } + + // Encode + encoded := EncodeReparsePoint(original) + + // Decode + decoded, err := DecodeReparsePoint(encoded) + if err != nil { + t.Fatalf("Failed to decode: %v", err) + } + + // Verify + if decoded.Target != original.Target { + t.Errorf("Target mismatch: got %q, want %q", decoded.Target, original.Target) + } + if !decoded.IsLxSymlink { + t.Errorf("IsLxSymlink should be true") + } +} + +func TestLxSymlinkRelativePath(t *testing.T) { + // Test LX symlink with relative path + original := &ReparsePoint{ + Target: testLxSymlinkRelativePath, + IsMountPoint: false, + IsLxSymlink: true, + } + + // Encode + encoded := EncodeReparsePoint(original) + + // Decode + decoded, err := DecodeReparsePoint(encoded) + if err != nil { + t.Fatalf("Failed to decode: %v", err) + } + + // Verify + if decoded.Target != original.Target { + t.Errorf("Target mismatch: got %q, want %q", decoded.Target, original.Target) + } + if !decoded.IsLxSymlink { + t.Errorf("IsLxSymlink should be true") + } +} + +func TestLxSymlinkSpecialCharacters(t *testing.T) { + // Test LX symlink with special characters and Unicode + original := &ReparsePoint{ + Target: testLxSymlinkSpecialCharsPath, + IsMountPoint: false, + IsLxSymlink: true, + } + + // Encode + encoded := EncodeReparsePoint(original) + + // Decode + decoded, err := DecodeReparsePoint(encoded) + if err != nil { + t.Fatalf("Failed to decode: %v", err) + } + + // Verify + if decoded.Target != original.Target { + t.Errorf("Target mismatch: got %q, want %q", decoded.Target, original.Target) + } + if !decoded.IsLxSymlink { + t.Errorf("IsLxSymlink should be true") + } +} + +func TestEncodeReparsePointNil(t *testing.T) { + // Test encoding a nil ReparsePoint + encoded := EncodeReparsePoint(nil) + if encoded != nil { + t.Errorf("Expected nil result for nil input, got %v", encoded) + } +}