Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
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
141 changes: 141 additions & 0 deletions compiler/air/embed_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,141 @@
package air

import (
"crypto/sha256"
"encoding/hex"
"fmt"
"os"
"path/filepath"
"strings"
"testing"

"github.com/akonwi/ard/checker"
"github.com/akonwi/ard/parse"
)

func TestLowerEmbeddedExactFilesIntoBlobReferences(t *testing.T) {
root := t.TempDir()
if err := os.WriteFile(filepath.Join(root, "ard.toml"), []byte("name = \"app\"\nard = \">= 0.1.0\"\n"), 0o644); err != nil {
t.Fatal(err)
}
content := []byte("shared embedded contents\n")
if err := os.WriteFile(filepath.Join(root, "asset.txt"), content, 0o644); err != nil {
t.Fatal(err)
}
mainPath := filepath.Join(root, "main.ard")
parsed := parse.Parse([]byte("use ard/embed\nlet page = embed::text(\"asset.txt\")\nlet raw = embed::bytes(\"asset.txt\")\nlet files = embed::fs([\"asset.txt\"])\n"), mainPath)
if len(parsed.Errors) > 0 {
t.Fatalf("parse errors: %v", parsed.Errors)
}
resolver, err := checker.NewModuleResolver(root)
if err != nil {
t.Fatal(err)
}
checked := checker.New(mainPath, parsed.Program, resolver)
checked.Check()
if checked.HasErrors() {
t.Fatalf("checker diagnostics: %#v", checked.Diagnostics())
}
if err := os.Remove(filepath.Join(root, "asset.txt")); err != nil {
t.Fatal(err)
}

program, err := Lower(checked.Module())
if err != nil {
t.Fatalf("Lower: %v", err)
}
if len(program.EmbeddedBlobs) != 1 {
t.Fatalf("embedded blobs = %d, want 1", len(program.EmbeddedBlobs))
}
if string(program.EmbeddedBlobs[0].Data) != string(content) {
t.Fatalf("blob data = %q", program.EmbeddedBlobs[0].Data)
}
if len(program.Globals) != 3 {
t.Fatalf("globals = %d, want 3", len(program.Globals))
}
text := program.Globals[0].Initializer.Value
bytes := program.Globals[1].Initializer.Value
if text.Kind != ExprEmbeddedText || bytes.Kind != ExprEmbeddedBytes {
t.Fatalf("embedded kinds = %d, %d", text.Kind, bytes.Kind)
}
if text.EmbeddedBlobPayload().Blob != 0 || bytes.EmbeddedBlobPayload().Blob != 0 {
t.Fatalf("embedded blob refs = %d, %d", text.EmbeddedBlobPayload().Blob, bytes.EmbeddedBlobPayload().Blob)
}
files := program.Globals[2].Initializer.Value
if files.Kind != ExprMakeEmbeddedFS || files.EmbeddedSetPayload().Set != 0 {
t.Fatalf("embedded filesystem expression = %#v", files)
}
if len(program.EmbeddedSets) != 1 || len(program.EmbeddedSets[0].Entries) != 1 || program.EmbeddedSets[0].Entries[0].Path != "asset.txt" {
t.Fatalf("embedded sets = %#v", program.EmbeddedSets)
}

encoded, err := SerializeProgram(program)
if err != nil {
t.Fatalf("SerializeProgram: %v", err)
}
roundTrip, err := DeserializeProgram(encoded)
if err != nil {
t.Fatalf("DeserializeProgram: %v", err)
}
if len(roundTrip.EmbeddedBlobs) != 1 || string(roundTrip.EmbeddedBlobs[0].Data) != string(content) {
t.Fatalf("round-trip embedded blobs = %#v", roundTrip.EmbeddedBlobs)
}
if len(roundTrip.EmbeddedSets) != 1 || roundTrip.EmbeddedSets[0].Digest != program.EmbeddedSets[0].Digest {
t.Fatalf("round-trip embedded sets = %#v", roundTrip.EmbeddedSets)
}

malformed := *program
duplicate := program.EmbeddedBlobs[0]
duplicate.ID = 1
malformed.EmbeddedBlobs = append(append([]EmbeddedBlob(nil), program.EmbeddedBlobs...), duplicate)
if err := Validate(&malformed); err == nil || !strings.Contains(err.Error(), "duplicates digest") {
t.Fatalf("duplicate embedded blob validation error = %v", err)
}

invalidPath := *program
invalidPath.EmbeddedSets = append([]EmbeddedSet(nil), program.EmbeddedSets...)
invalidPath.EmbeddedSets[0].Entries = append([]EmbeddedEntry(nil), program.EmbeddedSets[0].Entries...)
invalidPath.EmbeddedSets[0].Entries[0].Path = "../escape"
if err := Validate(&invalidPath); err == nil || !strings.Contains(err.Error(), "invalid path") {
t.Fatalf("invalid embedded set path validation error = %v", err)
}
}

func embeddedSetTestDigest(set EmbeddedSet, blobs []EmbeddedBlob) string {
hash := sha256.New()
_, _ = hash.Write([]byte(set.OwnerPackageIdentity))
_, _ = hash.Write([]byte{0})
for _, entry := range set.Entries {
_, _ = hash.Write([]byte(entry.Path))
_, _ = hash.Write([]byte{0})
_, _ = hash.Write([]byte(blobs[entry.Blob].Digest))
_, _ = hash.Write([]byte{0})
}
return hex.EncodeToString(hash.Sum(nil))
}

func TestValidateEmbeddedResourceLimits(t *testing.T) {
blobData := make([]byte, checker.MaxEmbeddedFileBytes)
blobSum := sha256.Sum256(blobData)
blob := EmbeddedBlob{ID: 0, Data: blobData, Digest: hex.EncodeToString(blobSum[:])}

oversizedSet := EmbeddedSet{ID: 0, OwnerPackageIdentity: "app"}
for index := 0; index < 5; index++ {
oversizedSet.Entries = append(oversizedSet.Entries, EmbeddedEntry{Path: fmt.Sprintf("%05d.bin", index), Blob: 0})
}
oversizedSet.Digest = embeddedSetTestDigest(oversizedSet, []EmbeddedBlob{blob})
if err := Validate(&Program{EmbeddedBlobs: []EmbeddedBlob{blob}, EmbeddedSets: []EmbeddedSet{oversizedSet}}); err == nil || !strings.Contains(err.Error(), "embedded set") {
t.Fatalf("oversized embedded set validation error = %v", err)
}

zeroSum := sha256.Sum256(nil)
zeroBlob := EmbeddedBlob{ID: 0, Digest: hex.EncodeToString(zeroSum[:])}
tooMany := EmbeddedSet{ID: 0, OwnerPackageIdentity: "app"}
for index := 0; index <= checker.MaxEmbeddedProgramFileCount; index++ {
tooMany.Entries = append(tooMany.Entries, EmbeddedEntry{Path: fmt.Sprintf("%05d.txt", index), Blob: 0})
}
tooMany.Digest = embeddedSetTestDigest(tooMany, []EmbeddedBlob{zeroBlob})
if err := Validate(&Program{EmbeddedBlobs: []EmbeddedBlob{zeroBlob}, EmbeddedSets: []EmbeddedSet{tooMany}}); err == nil || !strings.Contains(err.Error(), "program limit") {
t.Fatalf("embedded file-count validation error = %v", err)
}
}
10 changes: 10 additions & 0 deletions compiler/air/expr_payload_accessors.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,16 @@ func (e Expr) TextPayload() *TextExprPayload {
return payload
}

func (e Expr) EmbeddedBlobPayload() *EmbeddedBlobExprPayload {
payload, _ := e.Payload.(*EmbeddedBlobExprPayload)
return payload
}

func (e Expr) EmbeddedSetPayload() *EmbeddedSetExprPayload {
payload, _ := e.Payload.(*EmbeddedSetExprPayload)
return payload
}

func (e Expr) BoolPayload() *BoolExprPayload {
payload, _ := e.Payload.(*BoolExprPayload)
return payload
Expand Down
16 changes: 16 additions & 0 deletions compiler/air/expr_payloads.go
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,18 @@ type TextExprPayload struct {

func (*TextExprPayload) exprPayload() {}

type EmbeddedBlobExprPayload struct {
Blob EmbeddedBlobID
}

func (*EmbeddedBlobExprPayload) exprPayload() {}

type EmbeddedSetExprPayload struct {
Set EmbeddedSetID
}

func (*EmbeddedSetExprPayload) exprPayload() {}

type BoolExprPayload struct {
Value bool
}
Expand Down Expand Up @@ -242,6 +254,10 @@ func exprPayloadIsTypedNil(payload ExprPayload) bool {
switch payload := payload.(type) {
case *TextExprPayload:
return payload == nil
case *EmbeddedBlobExprPayload:
return payload == nil
case *EmbeddedSetExprPayload:
return payload == nil
case *BoolExprPayload:
return payload == nil
case *EnumExprPayload:
Expand Down
128 changes: 115 additions & 13 deletions compiler/air/lower.go
Original file line number Diff line number Diff line change
@@ -1,6 +1,9 @@
package air

import (
"bytes"
"crypto/sha256"
"encoding/hex"
"fmt"
"sort"
"strconv"
Expand Down Expand Up @@ -58,13 +61,15 @@ func LowerModulesWithOptions(modules []checker.Module, options LowerOptions) (*P
type lowerer struct {
program Program

moduleByPath map[string]ModuleID
moduleByName map[string]checker.Module
typeInterner *typeInterner
traits map[string]TraitID
impls map[string]ImplID
functions map[string]FunctionID
globals map[string]GlobalID
moduleByPath map[string]ModuleID
moduleByName map[string]checker.Module
typeInterner *typeInterner
traits map[string]TraitID
impls map[string]ImplID
functions map[string]FunctionID
globals map[string]GlobalID
embeddedBlobs map[string]EmbeddedBlobID
embeddedSets map[string]EmbeddedSetID

cacheMethodLookups bool
structMethodsByOwner map[checker.MethodOwner]map[string]*checker.FunctionDef
Expand Down Expand Up @@ -110,12 +115,14 @@ func newLowerer(options LowerOptions, rootCount int) *lowerer {
Entry: NoFunction,
Script: NoFunction,
},
moduleByPath: map[string]ModuleID{},
moduleByName: map[string]checker.Module{},
traits: map[string]TraitID{},
impls: map[string]ImplID{},
functions: map[string]FunctionID{},
globals: map[string]GlobalID{},
moduleByPath: map[string]ModuleID{},
moduleByName: map[string]checker.Module{},
traits: map[string]TraitID{},
impls: map[string]ImplID{},
functions: map[string]FunctionID{},
globals: map[string]GlobalID{},
embeddedBlobs: map[string]EmbeddedBlobID{},
embeddedSets: map[string]EmbeddedSetID{},

cacheMethodLookups: rootCount == 1,
unresolvedTypeVarByType: map[checker.Type]bool{},
Expand Down Expand Up @@ -172,6 +179,60 @@ func (l *lowerer) functionHasUnresolvedTypeVar(def *checker.FunctionDef) bool {
return l.typeHasUnresolvedTypeVar(def)
}

func (l *lowerer) internEmbeddedBlob(data []byte, direct bool) (EmbeddedBlobID, error) {
sum := sha256.Sum256(data)
digest := hex.EncodeToString(sum[:])
if id, ok := l.embeddedBlobs[digest]; ok {
if !bytes.Equal(l.program.EmbeddedBlobs[id].Data, data) {
return 0, fmt.Errorf("embedded resource digest collision for %s", digest)
}
if direct {
l.program.EmbeddedBlobs[id].Direct = true
}
return id, nil
}
id := EmbeddedBlobID(len(l.program.EmbeddedBlobs))
l.program.EmbeddedBlobs = append(l.program.EmbeddedBlobs, EmbeddedBlob{
ID: id,
Data: append([]byte(nil), data...),
Digest: digest,
Direct: direct,
})
l.embeddedBlobs[digest] = id
return id, nil
}

func (l *lowerer) internEmbeddedSet(set checker.EmbeddedFileSet) (EmbeddedSetID, error) {
entries := make([]EmbeddedEntry, len(set.Entries))
hash := sha256.New()
_, _ = hash.Write([]byte(set.OwnerPackageIdentity))
_, _ = hash.Write([]byte{0})
for index, entry := range set.Entries {
blob, err := l.internEmbeddedBlob(entry.Data, false)
if err != nil {
return 0, err
}
entries[index] = EmbeddedEntry{Path: entry.LogicalPath, Blob: blob}
_, _ = hash.Write([]byte(entry.LogicalPath))
_, _ = hash.Write([]byte{0})
_, _ = hash.Write([]byte(l.program.EmbeddedBlobs[blob].Digest))
_, _ = hash.Write([]byte{0})
}
digest := hex.EncodeToString(hash.Sum(nil))
if id, ok := l.embeddedSets[digest]; ok {
return id, nil
}
id := EmbeddedSetID(len(l.program.EmbeddedSets))
l.program.EmbeddedSets = append(l.program.EmbeddedSets, EmbeddedSet{
ID: id,
OwnerPackageIdentity: set.OwnerPackageIdentity,
Entries: entries,
Digest: digest,
})
l.embeddedSets[digest] = id
return id, nil
}

func (l *lowerer) mustIntern(t checker.Type) TypeID {
id, err := l.internType(t)
if err != nil {
Expand Down Expand Up @@ -2485,6 +2546,8 @@ func (l *lowerer) internAtomicOrTraitType(t checker.Type) (TypeID, error) {
info.Kind = TypeRune
case checker.Str:
info.Kind = TypeStr
case checker.EmbeddedFS:
info.Kind = TypeEmbeddedFS
case checker.Any:
info.Kind = TypeAny
default:
Expand Down Expand Up @@ -3929,6 +3992,24 @@ func (fl *functionLowerer) lowerExpr(expr checker.Expression) (*Expr, error) {
return &Expr{Kind: ExprConstBool, Type: typeID, Payload: &BoolExprPayload{Value: e.Value}}, nil
case *checker.StrLiteral:
return &Expr{Kind: ExprConstStr, Type: typeID, Payload: &TextExprPayload{Value: e.Value}}, nil
case *checker.EmbeddedText:
blob, err := fl.l.internEmbeddedBlob(e.Resource.Data, true)
if err != nil {
return nil, err
}
return &Expr{Kind: ExprEmbeddedText, Type: typeID, Payload: &EmbeddedBlobExprPayload{Blob: blob}}, nil
case *checker.EmbeddedBytes:
blob, err := fl.l.internEmbeddedBlob(e.Resource.Data, true)
if err != nil {
return nil, err
}
return &Expr{Kind: ExprEmbeddedBytes, Type: typeID, Payload: &EmbeddedBlobExprPayload{Blob: blob}}, nil
case *checker.EmbeddedFSValue:
set, err := fl.l.internEmbeddedSet(e.Set)
if err != nil {
return nil, err
}
return &Expr{Kind: ExprMakeEmbeddedFS, Type: typeID, Payload: &EmbeddedSetExprPayload{Set: set}}, nil
case *checker.RuneLiteral:
return &Expr{Kind: ExprConstInt, Type: typeID, Payload: &TextExprPayload{Value: strconv.Itoa(int(e.Value))}}, nil
case *checker.NeverCoercion:
Expand Down Expand Up @@ -4270,6 +4351,27 @@ func (fl *functionLowerer) lowerExpr(expr checker.Expression) (*Expr, error) {
return fl.lowerInstanceMethod(typeID, e)
case *checker.StrMethod:
return fl.lowerStrMethod(typeID, e)
case *checker.EmbeddedFSMethod:
target, err := fl.lowerExpr(e.Subject)
if err != nil {
return nil, err
}
args, err := fl.lowerArgs(e.Args)
if err != nil {
return nil, err
}
kind := ExprEmbeddedFSReadFile
switch e.Kind {
case checker.EmbeddedFSReadText:
kind = ExprEmbeddedFSReadText
case checker.EmbeddedFSReadDir:
kind = ExprEmbeddedFSReadDir
case checker.EmbeddedFSStat:
kind = ExprEmbeddedFSStat
case checker.EmbeddedFSSub:
kind = ExprEmbeddedFSSub
}
return &Expr{Kind: kind, Type: typeID, Target: target, Args: args}, nil
case *checker.ByteMethod:
if e.Kind == checker.ByteToInt {
return fl.lowerUnary(ExprToInt, typeID, e.Subject)
Expand Down
8 changes: 8 additions & 0 deletions compiler/air/nodes.go
Original file line number Diff line number Diff line change
Expand Up @@ -62,6 +62,14 @@ const (
ExprConstFloat
ExprConstBool
ExprConstStr
ExprEmbeddedText
ExprEmbeddedBytes
ExprMakeEmbeddedFS
ExprEmbeddedFSReadFile
ExprEmbeddedFSReadText
ExprEmbeddedFSReadDir
ExprEmbeddedFSStat
ExprEmbeddedFSSub
ExprPanic
ExprLoadLocal
ExprLoadGlobal
Expand Down
2 changes: 2 additions & 0 deletions compiler/air/serialize.go
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,8 @@ import (

func init() {
gob.Register(&TextExprPayload{})
gob.Register(&EmbeddedBlobExprPayload{})
gob.Register(&EmbeddedSetExprPayload{})
gob.Register(&BoolExprPayload{})
gob.Register(&EnumExprPayload{})
gob.Register(&LocalExprPayload{})
Expand Down
Loading
Loading