diff options
| author | Ophestra <cat@gensokyo.uk> | 2026-03-17 15:35:58 +0900 |
|---|---|---|
| committer | Ophestra <cat@gensokyo.uk> | 2026-03-17 15:39:03 +0900 |
| commit | 6d015a949e27f20c86409c7d78053f0b0493f165 (patch) | |
| tree | 4d2bc1e9596fa30ea03e4c0cc575e6a7a969249d /check | |
| parent | e9a72490db44426605debd3cc756a7924931d706 (diff) | |
check: move from container
This package is not container specific, and widely used across the project.
Signed-off-by: Ophestra <cat@gensokyo.uk>
Diffstat (limited to 'check')
| -rw-r--r-- | check/absolute.go | 132 | ||||
| -rw-r--r-- | check/absolute_test.go | 401 | ||||
| -rw-r--r-- | check/overlay.go | 30 | ||||
| -rw-r--r-- | check/overlay_test.go | 31 |
4 files changed, 594 insertions, 0 deletions
diff --git a/check/absolute.go b/check/absolute.go new file mode 100644 index 00000000..d17d5e1a --- /dev/null +++ b/check/absolute.go @@ -0,0 +1,132 @@ +// Package check provides types yielding values checked to meet a condition. +package check + +import ( + "encoding/json" + "errors" + "fmt" + "path" + "slices" + "strings" + "syscall" + "unique" +) + +// AbsoluteError is returned by [NewAbs] and holds the invalid pathname. +type AbsoluteError string + +func (e AbsoluteError) Error() string { + return fmt.Sprintf("path %q is not absolute", string(e)) +} + +func (e AbsoluteError) Is(target error) bool { + var ce AbsoluteError + if !errors.As(target, &ce) { + return errors.Is(target, syscall.EINVAL) + } + return e == ce +} + +// Absolute holds a pathname checked to be absolute. +type Absolute struct{ pathname unique.Handle[string] } + +// ok returns whether [Absolute] is not the zero value. +func (a *Absolute) ok() bool { return a != nil && *a != (Absolute{}) } + +// unsafeAbs returns [check.Absolute] on any string value. +func unsafeAbs(pathname string) *Absolute { + return &Absolute{unique.Make(pathname)} +} + +// String returns the checked pathname. +func (a *Absolute) String() string { + if !a.ok() { + panic("attempted use of zero Absolute") + } + return a.pathname.Value() +} + +// Handle returns the underlying [unique.Handle]. +func (a *Absolute) Handle() unique.Handle[string] { + return a.pathname +} + +// Is efficiently compares the underlying pathname. +func (a *Absolute) Is(v *Absolute) bool { + if a == nil && v == nil { + return true + } + return a.ok() && v.ok() && a.pathname == v.pathname +} + +// NewAbs checks pathname and returns a new [Absolute] if pathname is absolute. +func NewAbs(pathname string) (*Absolute, error) { + if !path.IsAbs(pathname) { + return nil, AbsoluteError(pathname) + } + return unsafeAbs(pathname), nil +} + +// MustAbs calls [NewAbs] and panics on error. +func MustAbs(pathname string) *Absolute { + if a, err := NewAbs(pathname); err != nil { + panic(err) + } else { + return a + } +} + +// Append calls [path.Join] with [Absolute] as the first element. +func (a *Absolute) Append(elem ...string) *Absolute { + return unsafeAbs(path.Join(append([]string{a.String()}, elem...)...)) +} + +// Dir calls [path.Dir] with [Absolute] as its argument. +func (a *Absolute) Dir() *Absolute { return unsafeAbs(path.Dir(a.String())) } + +// GobEncode returns the checked pathname. +func (a *Absolute) GobEncode() ([]byte, error) { + return []byte(a.String()), nil +} + +// GobDecode stores data if it represents an absolute pathname. +func (a *Absolute) GobDecode(data []byte) error { + pathname := string(data) + if !path.IsAbs(pathname) { + return AbsoluteError(pathname) + } + a.pathname = unique.Make(pathname) + return nil +} + +// MarshalJSON returns a JSON representation of the checked pathname. +func (a *Absolute) MarshalJSON() ([]byte, error) { + return json.Marshal(a.String()) +} + +// UnmarshalJSON stores data if it represents an absolute pathname. +func (a *Absolute) UnmarshalJSON(data []byte) error { + var pathname string + if err := json.Unmarshal(data, &pathname); err != nil { + return err + } + if !path.IsAbs(pathname) { + return AbsoluteError(pathname) + } + a.pathname = unique.Make(pathname) + return nil +} + +// SortAbs calls [slices.SortFunc] for a slice of [Absolute]. +func SortAbs(x []*Absolute) { + slices.SortFunc(x, func(a, b *Absolute) int { + return strings.Compare(a.String(), b.String()) + }) +} + +// CompactAbs calls [slices.CompactFunc] for a slice of [Absolute]. +func CompactAbs(s []*Absolute) []*Absolute { + return slices.CompactFunc(s, func(a *Absolute, b *Absolute) bool { + return a.Is(b) + }) +} diff --git a/check/absolute_test.go b/check/absolute_test.go new file mode 100644 index 00000000..445ff2e2 --- /dev/null +++ b/check/absolute_test.go @@ -0,0 +1,401 @@ +package check_test + +import ( + "bytes" + "encoding/gob" + "encoding/json" + "errors" + "reflect" + "strings" + "syscall" + "testing" + _ "unsafe" // for go:linkname + + . "hakurei.app/check" +) + +// unsafeAbs returns check.Absolute on any string value. +// +//go:linkname unsafeAbs hakurei.app/check.unsafeAbs +func unsafeAbs(pathname string) *Absolute + +func TestAbsoluteError(t *testing.T) { + t.Parallel() + + testCases := []struct { + name string + + err error + cmp error + ok bool + }{ + {"EINVAL", new(AbsoluteError), syscall.EINVAL, true}, + {"not EINVAL", new(AbsoluteError), syscall.EBADE, false}, + {"ne val", new(AbsoluteError), AbsoluteError("etc"), false}, + {"equals", AbsoluteError("etc"), AbsoluteError("etc"), true}, + } + + for _, tc := range testCases { + if got := errors.Is(tc.err, tc.cmp); got != tc.ok { + t.Errorf("Is: %v, want %v", got, tc.ok) + } + } + + t.Run("string", func(t *testing.T) { + t.Parallel() + + want := `path "etc" is not absolute` + if got := (AbsoluteError("etc")).Error(); got != want { + t.Errorf("Error: %q, want %q", got, want) + } + }) +} + +func TestNewAbs(t *testing.T) { + t.Parallel() + + testCases := []struct { + name string + + pathname string + want *Absolute + wantErr error + }{ + {"good", "/etc", MustAbs("/etc"), nil}, + {"not absolute", "etc", nil, AbsoluteError("etc")}, + {"zero", "", nil, AbsoluteError("")}, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + got, err := NewAbs(tc.pathname) + if !reflect.DeepEqual(got, tc.want) { + t.Errorf("NewAbs: %#v, want %#v", got, tc.want) + } + if !errors.Is(err, tc.wantErr) { + t.Errorf("NewAbs: error = %v, want %v", err, tc.wantErr) + } + }) + } + + t.Run("must", func(t *testing.T) { + t.Parallel() + + defer func() { + wantPanic := AbsoluteError("etc") + + if r := recover(); !reflect.DeepEqual(r, wantPanic) { + t.Errorf("MustAbs: panic = %v; want %v", r, wantPanic) + } + }() + + MustAbs("etc") + }) +} + +func TestAbsoluteString(t *testing.T) { + t.Run("passthrough", func(t *testing.T) { + t.Parallel() + + pathname := "/etc" + if got := unsafeAbs(pathname).String(); got != pathname { + t.Errorf("String: %q, want %q", got, pathname) + } + }) + + t.Run("zero", func(t *testing.T) { + t.Parallel() + + defer func() { + wantPanic := "attempted use of zero Absolute" + + if r := recover(); r != wantPanic { + t.Errorf("String: panic = %v, want %v", r, wantPanic) + } + }() + + panic(new(Absolute).String()) + }) +} + +func TestAbsoluteIs(t *testing.T) { + t.Parallel() + + testCases := []struct { + name string + a, v *Absolute + want bool + }{ + {"nil", (*Absolute)(nil), (*Absolute)(nil), true}, + {"nil a", (*Absolute)(nil), MustAbs("/"), false}, + {"nil v", MustAbs("/"), (*Absolute)(nil), false}, + {"zero", new(Absolute), new(Absolute), false}, + {"zero a", new(Absolute), MustAbs("/"), false}, + {"zero v", MustAbs("/"), new(Absolute), false}, + {"equals", MustAbs("/"), MustAbs("/"), true}, + } + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + if got := tc.a.Is(tc.v); got != tc.want { + t.Errorf("Is: %v, want %v", got, tc.want) + } + }) + } +} + +type sCheck struct { + Pathname *Absolute `json:"val"` + Magic uint64 `json:"magic"` +} + +func TestCodecAbsolute(t *testing.T) { + t.Parallel() + + testCases := []struct { + name string + a *Absolute + + wantErr error + + gob, sGob string + json, sJson string + }{ + {"nil", nil, nil, + "\x00", "\x00", + `null`, `{"val":null,"magic":3236757504}`}, + + {"good", MustAbs("/etc"), + nil, + "\t\x7f\x05\x01\x02\xff\x82\x00\x00\x00\b\xff\x80\x00\x04/etc", + ",\xff\x83\x03\x01\x01\x06sCheck\x01\xff\x84\x00\x01\x02\x01\bPathname\x01\xff\x80\x00\x01\x05Magic\x01\x06\x00\x00\x00\t\x7f\x05\x01\x02\xff\x82\x00\x00\x00\x0f\xff\x84\x01\x04/etc\x01\xfc\xc0\xed\x00\x00\x00", + + `"/etc"`, `{"val":"/etc","magic":3236757504}`}, + {"not absolute", nil, + AbsoluteError("etc"), + "\t\x7f\x05\x01\x02\xff\x82\x00\x00\x00\a\xff\x80\x00\x03etc", + ",\xff\x83\x03\x01\x01\x06sCheck\x01\xff\x84\x00\x01\x02\x01\bPathname\x01\xff\x80\x00\x01\x05Magic\x01\x06\x00\x00\x00\t\x7f\x05\x01\x02\xff\x82\x00\x00\x00\x0f\xff\x84\x01\x03etc\x01\xfb\x01\x81\xda\x00\x00\x00", + + `"etc"`, `{"val":"etc","magic":3236757504}`}, + {"zero", nil, + new(AbsoluteError), + "\t\x7f\x05\x01\x02\xff\x82\x00\x00\x00\x04\xff\x80\x00\x00", + ",\xff\x83\x03\x01\x01\x06sCheck\x01\xff\x84\x00\x01\x02\x01\bPathname\x01\xff\x80\x00\x01\x05Magic\x01\x06\x00\x00\x00\t\x7f\x05\x01\x02\xff\x82\x00\x00\x00\f\xff\x84\x01\x00\x01\xfb\x01\x81\xda\x00\x00\x00", + `""`, `{"val":"","magic":3236757504}`}, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + t.Run("gob", func(t *testing.T) { + if tc.gob == "\x00" && tc.sGob == "\x00" { + // these values mark the current test to skip gob + return + } + t.Parallel() + + t.Run("encode", func(t *testing.T) { + t.Parallel() + + // encode is unchecked + if errors.Is(tc.wantErr, syscall.EINVAL) { + return + } + + { + buf := new(bytes.Buffer) + err := gob.NewEncoder(buf).Encode(tc.a) + if !errors.Is(err, tc.wantErr) { + t.Errorf("Encode: error = %v, want %v", err, tc.wantErr) + } + if tc.wantErr != nil { + goto checkSEncode + } + if buf.String() != tc.gob { + t.Errorf("Encode:\n%q\nwant:\n%q", buf.String(), tc.gob) + } + } + + checkSEncode: + { + buf := new(bytes.Buffer) + err := gob.NewEncoder(buf).Encode(&sCheck{tc.a, syscall.MS_MGC_VAL}) + if !errors.Is(err, tc.wantErr) { + t.Errorf("Encode: error = %v, want %v", err, tc.wantErr) + } + if tc.wantErr != nil { + return + } + if buf.String() != tc.sGob { + t.Errorf("Encode:\n%q\nwant:\n%q", buf.String(), tc.sGob) + } + } + }) + + t.Run("decode", func(t *testing.T) { + t.Parallel() + + { + var gotA *Absolute + err := gob.NewDecoder(strings.NewReader(tc.gob)).Decode(&gotA) + if !errors.Is(err, tc.wantErr) { + t.Errorf("Decode: error = %v, want %v", err, tc.wantErr) + } + if tc.wantErr != nil { + goto checkSDecode + } + if !reflect.DeepEqual(tc.a, gotA) { + t.Errorf("Decode: %#v, want %#v", tc.a, gotA) + } + } + + checkSDecode: + { + var gotSCheck sCheck + err := gob.NewDecoder(strings.NewReader(tc.sGob)).Decode(&gotSCheck) + if !errors.Is(err, tc.wantErr) { + t.Errorf("Decode: error = %v, want %v", err, tc.wantErr) + } + if tc.wantErr != nil { + return + } + want := sCheck{tc.a, syscall.MS_MGC_VAL} + if !reflect.DeepEqual(gotSCheck, want) { + t.Errorf("Decode: %#v, want %#v", gotSCheck, want) + } + } + }) + + }) + + t.Run("json", func(t *testing.T) { + t.Parallel() + + t.Run("marshal", func(t *testing.T) { + t.Parallel() + + // marshal is unchecked + if errors.Is(tc.wantErr, syscall.EINVAL) { + return + } + + { + d, err := json.Marshal(tc.a) + if !errors.Is(err, tc.wantErr) { + t.Errorf("Marshal: error = %v, want %v", err, tc.wantErr) + } + if tc.wantErr != nil { + goto checkSMarshal + } + if string(d) != tc.json { + t.Errorf("Marshal:\n%s\nwant:\n%s", string(d), tc.json) + } + } + + checkSMarshal: + { + d, err := json.Marshal(&sCheck{tc.a, syscall.MS_MGC_VAL}) + if !errors.Is(err, tc.wantErr) { + t.Errorf("Marshal: error = %v, want %v", err, tc.wantErr) + } + if tc.wantErr != nil { + return + } + if string(d) != tc.sJson { + t.Errorf("Marshal:\n%s\nwant:\n%s", string(d), tc.sJson) + } + } + }) + + t.Run("unmarshal", func(t *testing.T) { + t.Parallel() + + { + var gotA *Absolute + err := json.Unmarshal([]byte(tc.json), &gotA) + if !errors.Is(err, tc.wantErr) { + t.Errorf("Unmarshal: error = %v, want %v", err, tc.wantErr) + } + if tc.wantErr != nil { + goto checkSUnmarshal + } + if !reflect.DeepEqual(tc.a, gotA) { + t.Errorf("Unmarshal: %#v, want %#v", tc.a, gotA) + } + } + + checkSUnmarshal: + { + var gotSCheck sCheck + err := json.Unmarshal([]byte(tc.sJson), &gotSCheck) + if !errors.Is(err, tc.wantErr) { + t.Errorf("Unmarshal: error = %v, want %v", err, tc.wantErr) + } + if tc.wantErr != nil { + return + } + want := sCheck{tc.a, syscall.MS_MGC_VAL} + if !reflect.DeepEqual(gotSCheck, want) { + t.Errorf("Unmarshal: %#v, want %#v", gotSCheck, want) + } + } + }) + }) + }) + } + + t.Run("json passthrough", func(t *testing.T) { + t.Parallel() + + wantErr := "invalid character ':' looking for beginning of value" + if err := new(Absolute).UnmarshalJSON([]byte(":3")); err == nil || err.Error() != wantErr { + t.Errorf("UnmarshalJSON: error = %v, want %s", err, wantErr) + } + }) +} + +func TestAbsoluteWrap(t *testing.T) { + t.Parallel() + + t.Run("join", func(t *testing.T) { + t.Parallel() + + want := "/etc/nix/nix.conf" + if got := MustAbs("/etc").Append("nix", "nix.conf"); got.String() != want { + t.Errorf("Append: %q, want %q", got, want) + } + }) + + t.Run("dir", func(t *testing.T) { + t.Parallel() + + want := "/" + if got := MustAbs("/etc").Dir(); got.String() != want { + t.Errorf("Dir: %q, want %q", got, want) + } + }) + + t.Run("sort", func(t *testing.T) { + t.Parallel() + + want := []*Absolute{MustAbs("/etc"), MustAbs("/proc"), MustAbs("/sys")} + got := []*Absolute{MustAbs("/proc"), MustAbs("/sys"), MustAbs("/etc")} + SortAbs(got) + if !reflect.DeepEqual(got, want) { + t.Errorf("SortAbs: %#v, want %#v", got, want) + } + }) + + t.Run("compact", func(t *testing.T) { + t.Parallel() + + want := []*Absolute{MustAbs("/etc"), MustAbs("/proc"), MustAbs("/sys")} + if got := CompactAbs([]*Absolute{MustAbs("/etc"), MustAbs("/proc"), MustAbs("/proc"), MustAbs("/sys")}); !reflect.DeepEqual(got, want) { + t.Errorf("CompactAbs: %#v, want %#v", got, want) + } + }) +} diff --git a/check/overlay.go b/check/overlay.go new file mode 100644 index 00000000..fbbcd1e6 --- /dev/null +++ b/check/overlay.go @@ -0,0 +1,30 @@ +package check + +import "strings" + +const ( + // SpecialOverlayEscape is the escape string for overlay mount options. + SpecialOverlayEscape = `\` + // SpecialOverlayOption is the separator string between overlay mount options. + SpecialOverlayOption = "," + // SpecialOverlayPath is the separator string between overlay paths. + SpecialOverlayPath = ":" +) + +// EscapeOverlayDataSegment escapes a string for formatting into the data +// argument of an overlay mount system call. +func EscapeOverlayDataSegment(s string) string { + if s == "" { + return "" + } + + if f := strings.SplitN(s, "\x00", 2); len(f) > 0 { + s = f[0] + } + + return strings.NewReplacer( + SpecialOverlayEscape, SpecialOverlayEscape+SpecialOverlayEscape, + SpecialOverlayOption, SpecialOverlayEscape+SpecialOverlayOption, + SpecialOverlayPath, SpecialOverlayEscape+SpecialOverlayPath, + ).Replace(s) +} diff --git a/check/overlay_test.go b/check/overlay_test.go new file mode 100644 index 00000000..9b32dbcb --- /dev/null +++ b/check/overlay_test.go @@ -0,0 +1,31 @@ +package check_test + +import ( + "testing" + + "hakurei.app/check" +) + +func TestEscapeOverlayDataSegment(t *testing.T) { + t.Parallel() + + testCases := []struct { + name string + s string + want string + }{ + {"zero", "", ""}, + {"multi", `\\\:,:,\\\`, `\\\\\\\:\,\:\,\\\\\\`}, + {"bwrap", `/path :,\`, `/path \:\,\\`}, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + if got := check.EscapeOverlayDataSegment(tc.s); got != tc.want { + t.Errorf("escapeOverlayDataSegment: %s, want %s", got, tc.want) + } + }) + } +} |
