@@ -0,0 +1,215 @@
|
||||
package composeedit
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"gopkg.in/yaml.v3"
|
||||
)
|
||||
|
||||
type ParseResult struct {
|
||||
Value any `json:"value"`
|
||||
}
|
||||
|
||||
type Patch struct {
|
||||
Path []string `json:"path"`
|
||||
Value any `json:"value"`
|
||||
Delete bool `json:"delete"`
|
||||
}
|
||||
|
||||
func Parse(src string) (ParseResult, error) {
|
||||
var doc yaml.Node
|
||||
if err := yaml.Unmarshal([]byte(src), &doc); err != nil {
|
||||
return ParseResult{}, fmt.Errorf("yaml: %w", err)
|
||||
}
|
||||
if len(doc.Content) == 0 {
|
||||
return ParseResult{Value: map[string]any{}}, nil
|
||||
}
|
||||
var v any
|
||||
if err := doc.Content[0].Decode(&v); err != nil {
|
||||
return ParseResult{}, err
|
||||
}
|
||||
v = normalize(v)
|
||||
return ParseResult{Value: v}, nil
|
||||
}
|
||||
|
||||
func Apply(src string, p Patch) (string, error) {
|
||||
if len(p.Path) == 0 {
|
||||
return "", errors.New("path is required")
|
||||
}
|
||||
var doc yaml.Node
|
||||
if err := yaml.Unmarshal([]byte(src), &doc); err != nil {
|
||||
return "", fmt.Errorf("yaml: %w", err)
|
||||
}
|
||||
if len(doc.Content) == 0 {
|
||||
return "", errors.New("empty yaml document")
|
||||
}
|
||||
root := doc.Content[0]
|
||||
parent, last, err := walkParent(root, p.Path, !p.Delete)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if p.Delete {
|
||||
if err := deleteChild(parent, last); err != nil {
|
||||
return "", err
|
||||
}
|
||||
} else {
|
||||
n, err := valueNode(p.Value)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if err := setChild(parent, last, n); err != nil {
|
||||
return "", err
|
||||
}
|
||||
}
|
||||
var b bytes.Buffer
|
||||
enc := yaml.NewEncoder(&b)
|
||||
enc.SetIndent(2)
|
||||
if err := enc.Encode(&doc); err != nil {
|
||||
return "", err
|
||||
}
|
||||
_ = enc.Close()
|
||||
return strings.TrimSuffix(b.String(), "\n") + "\n", nil
|
||||
}
|
||||
|
||||
func normalize(v any) any {
|
||||
switch x := v.(type) {
|
||||
case map[string]any:
|
||||
out := make(map[string]any, len(x))
|
||||
for k, v := range x {
|
||||
out[k] = normalize(v)
|
||||
}
|
||||
return out
|
||||
case map[any]any:
|
||||
out := map[string]any{}
|
||||
for k, v := range x {
|
||||
out[fmt.Sprint(k)] = normalize(v)
|
||||
}
|
||||
return out
|
||||
case []any:
|
||||
out := make([]any, len(x))
|
||||
for i, v := range x {
|
||||
out[i] = normalize(v)
|
||||
}
|
||||
return out
|
||||
default:
|
||||
return x
|
||||
}
|
||||
}
|
||||
|
||||
func walkParent(root *yaml.Node, path []string, create bool) (*yaml.Node, string, error) {
|
||||
cur := root
|
||||
for _, seg := range path[:len(path)-1] {
|
||||
if cur.Kind == yaml.DocumentNode && len(cur.Content) > 0 {
|
||||
cur = cur.Content[0]
|
||||
}
|
||||
switch cur.Kind {
|
||||
case yaml.MappingNode:
|
||||
n := mapGet(cur, seg)
|
||||
if n == nil {
|
||||
if !create {
|
||||
return nil, "", fmt.Errorf("path %q not found", seg)
|
||||
}
|
||||
n = &yaml.Node{Kind: yaml.MappingNode, Tag: "!!map"}
|
||||
mapSet(cur, seg, n)
|
||||
}
|
||||
cur = n
|
||||
case yaml.SequenceNode:
|
||||
i, err := strconv.Atoi(seg)
|
||||
if err != nil || i < 0 || i >= len(cur.Content) {
|
||||
return nil, "", fmt.Errorf("invalid array index %q", seg)
|
||||
}
|
||||
cur = cur.Content[i]
|
||||
default:
|
||||
return nil, "", fmt.Errorf("cannot descend through scalar at %q", seg)
|
||||
}
|
||||
}
|
||||
return cur, path[len(path)-1], nil
|
||||
}
|
||||
func mapGet(m *yaml.Node, key string) *yaml.Node {
|
||||
for i := 0; i+1 < len(m.Content); i += 2 {
|
||||
if m.Content[i].Value == key {
|
||||
return m.Content[i+1]
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
func mapSet(m *yaml.Node, key string, v *yaml.Node) {
|
||||
for i := 0; i+1 < len(m.Content); i += 2 {
|
||||
if m.Content[i].Value == key {
|
||||
old := m.Content[i+1]
|
||||
v.HeadComment = old.HeadComment
|
||||
v.LineComment = old.LineComment
|
||||
v.FootComment = old.FootComment
|
||||
m.Content[i+1] = v
|
||||
return
|
||||
}
|
||||
}
|
||||
m.Content = append(m.Content, &yaml.Node{Kind: yaml.ScalarNode, Tag: "!!str", Value: key}, v)
|
||||
}
|
||||
func setChild(p *yaml.Node, key string, v *yaml.Node) error {
|
||||
switch p.Kind {
|
||||
case yaml.MappingNode:
|
||||
mapSet(p, key, v)
|
||||
return nil
|
||||
case yaml.SequenceNode:
|
||||
if key == "-" {
|
||||
p.Content = append(p.Content, v)
|
||||
return nil
|
||||
}
|
||||
i, err := strconv.Atoi(key)
|
||||
if err != nil || i < 0 || i > len(p.Content) {
|
||||
return fmt.Errorf("invalid array index %q", key)
|
||||
}
|
||||
if i == len(p.Content) {
|
||||
p.Content = append(p.Content, v)
|
||||
} else {
|
||||
old := p.Content[i]
|
||||
v.HeadComment, v.LineComment, v.FootComment = old.HeadComment, old.LineComment, old.FootComment
|
||||
p.Content[i] = v
|
||||
}
|
||||
return nil
|
||||
default:
|
||||
return errors.New("parent is not a map or array")
|
||||
}
|
||||
}
|
||||
func deleteChild(p *yaml.Node, key string) error {
|
||||
switch p.Kind {
|
||||
case yaml.MappingNode:
|
||||
for i := 0; i+1 < len(p.Content); i += 2 {
|
||||
if p.Content[i].Value == key {
|
||||
p.Content = append(p.Content[:i], p.Content[i+2:]...)
|
||||
return nil
|
||||
}
|
||||
}
|
||||
return nil
|
||||
case yaml.SequenceNode:
|
||||
i, err := strconv.Atoi(key)
|
||||
if err != nil || i < 0 || i >= len(p.Content) {
|
||||
return fmt.Errorf("invalid array index %q", key)
|
||||
}
|
||||
p.Content = append(p.Content[:i], p.Content[i+1:]...)
|
||||
return nil
|
||||
default:
|
||||
return errors.New("parent is not a map or array")
|
||||
}
|
||||
}
|
||||
func valueNode(v any) (*yaml.Node, error) {
|
||||
raw, err := json.Marshal(v)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var x any
|
||||
if err := json.Unmarshal(raw, &x); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var n yaml.Node
|
||||
if err := n.Encode(x); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &n, nil
|
||||
}
|
||||
@@ -0,0 +1,31 @@
|
||||
package composeedit
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestPatchPreservesUnknownAndComments(t *testing.T) {
|
||||
src := "# top\nservices:\n app:\n image: nginx:old # keep\n x-future:\n magic: true\n deploy:\n replicas: 2\n"
|
||||
out, e := Apply(src, Patch{Path: []string{"services", "app", "image"}, Value: "nginx:new"})
|
||||
if e != nil {
|
||||
t.Fatal(e)
|
||||
}
|
||||
if !strings.Contains(out, "x-future:") || !strings.Contains(out, "magic: true") || !strings.Contains(out, "replicas: 2") || !strings.Contains(out, "# top") {
|
||||
t.Fatalf("preservation failed:\n%s", out)
|
||||
}
|
||||
r, e := Parse(out)
|
||||
if e != nil || r.Value == nil {
|
||||
t.Fatalf("parse %v", e)
|
||||
}
|
||||
}
|
||||
func TestArrayPatch(t *testing.T) {
|
||||
src := "services:\n app:\n ports:\n - 8080:80\n"
|
||||
out, e := Apply(src, Patch{Path: []string{"services", "app", "ports", "0"}, Value: "9090:80"})
|
||||
if e != nil {
|
||||
t.Fatal(e)
|
||||
}
|
||||
if !strings.Contains(out, "9090:80") {
|
||||
t.Fatal(out)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user