diff --git a/cmd/relspec/assets.go b/cmd/relspec/assets.go new file mode 100644 index 0000000..062171b --- /dev/null +++ b/cmd/relspec/assets.go @@ -0,0 +1,214 @@ +package main + +import ( + "context" + "fmt" + "os" + "path/filepath" + + "github.com/spf13/cobra" + + "git.warky.dev/wdevs/relspecgo/pkg/assetloader" + "git.warky.dev/wdevs/relspecgo/pkg/pgsql" +) + +var ( + assetsDir string + assetsConn string + assetsIgnoreErrors bool +) + +var assetsCmd = &cobra.Command{ + Use: "assets", + Short: "Load and execute asset manifests against a database", + Long: `Load local binary and text asset files into a PostgreSQL database. + +Assets are described by YAML manifests (assets.yaml) colocated with the files. +Each manifest entry specifies the file to load and the SQL call to execute. +File bytes are bound as native pgx parameters — never as SQL text literals — +so binary files stay byte-exact with no size or encoding limitations. + +Manifests must live in directories that follow the naming pattern used by +relspec scripts: + {priority}_{sequence}_{name}/ or {priority}-{sequence}-{name}/ + +This allows asset-loading steps to be ordered correctly alongside SQL scripts +in a migrate-apply pipeline. + +Manifest format (assets.yaml): + - file: invoice.md + call: | + INSERT INTO org.filepointer (rid_owner, filename, contenttype, jsonstore) + VALUES (1, :filename, 'text/markdown', jsonb_build_object('content', :bytes::text)) + - file: logo.png + call: UPDATE branding SET logo = :bytes WHERE id = 1 + params: + owner_id: "42" + +Built-in placeholders: + :bytes — the file's raw content as bytea + :filename — the base name of the file (string) + :any_key — a static value declared in the entry's params map`, +} + +var assetsListCmd = &cobra.Command{ + Use: "list", + Short: "List asset manifests from a directory", + Long: `List all asset manifest entries from a directory in execution order. + +The directory is scanned recursively for assets.yaml files located in +directories that follow the {priority}_{sequence}_{name} naming convention. + +Example: + relspec assets list --dir ./sql`, + RunE: runAssetsList, +} + +var assetsExecuteCmd = &cobra.Command{ + Use: "execute", + Short: "Execute asset manifests against a database", + Long: `Execute asset manifest entries from a directory against a PostgreSQL database. + +Asset manifests are executed in order: Priority (ascending), Sequence (ascending), +Directory name (alphabetical). By default, execution stops on the first error. +Use --ignore-errors to continue even when individual entries fail. + +PostgreSQL Connection String Examples: + postgres://username:password@localhost:5432/database_name + postgresql://user:pass@host/dbname?sslmode=disable + +Examples: + relspec assets execute --dir ./sql \ + --conn "postgres://user:pass@localhost:5432/mydb" + + relspec assets execute --dir ./sql \ + --conn "postgres://localhost/mydb" \ + --ignore-errors`, + RunE: runAssetsExecute, +} + +func init() { + assetsListCmd.Flags().StringVar(&assetsDir, "dir", "", "Directory to scan for asset manifests (required)") + if err := assetsListCmd.MarkFlagRequired("dir"); err != nil { + fmt.Fprintf(os.Stderr, "Error marking dir flag as required: %v\n", err) + } + + assetsExecuteCmd.Flags().StringVar(&assetsDir, "dir", "", "Directory to scan for asset manifests (required)") + assetsExecuteCmd.Flags().StringVar(&assetsConn, "conn", "", "PostgreSQL connection string (required)") + assetsExecuteCmd.Flags().BoolVar(&assetsIgnoreErrors, "ignore-errors", false, "Continue executing even if entries fail") + if err := assetsExecuteCmd.MarkFlagRequired("dir"); err != nil { + fmt.Fprintf(os.Stderr, "Error marking dir flag as required: %v\n", err) + } + if err := assetsExecuteCmd.MarkFlagRequired("conn"); err != nil { + fmt.Fprintf(os.Stderr, "Error marking conn flag as required: %v\n", err) + } + + assetsCmd.AddCommand(assetsListCmd) + assetsCmd.AddCommand(assetsExecuteCmd) +} + +func runAssetsList(cmd *cobra.Command, args []string) error { + fmt.Fprintf(os.Stderr, "\n=== Asset Manifests List ===\n") + fmt.Fprintf(os.Stderr, "Directory: %s\n\n", assetsDir) + + items, err := assetloader.ScanDir(assetsDir) + if err != nil { + return fmt.Errorf("scanning directory: %w", err) + } + + if len(items) == 0 { + fmt.Fprintf(os.Stderr, "No asset manifests found.\n\n") + return nil + } + + fmt.Fprintf(os.Stderr, "Found %d asset entry(ies) in execution order:\n\n", len(items)) + fmt.Fprintf(os.Stderr, "%-4s %-10s %-8s %-20s %s\n", "No.", "Priority", "Sequence", "Dir", "File") + fmt.Fprintf(os.Stderr, "%-4s %-10s %-8s %-20s %s\n", "----", "--------", "--------", "--------------------", "----") + + for i, item := range items { + fmt.Fprintf(os.Stderr, "%-4d %-10d %-8d %-20s %s\n", + i+1, + item.Priority, + item.Sequence, + item.DirName, + filepath.Base(item.Entry.File), + ) + } + + fmt.Fprintf(os.Stderr, "\n") + return nil +} + +func runAssetsExecute(cmd *cobra.Command, args []string) error { + fmt.Fprintf(os.Stderr, "\n=== Asset Manifests Execution ===\n") + fmt.Fprintf(os.Stderr, "Started at: %s\n", getCurrentTimestamp()) + fmt.Fprintf(os.Stderr, "Directory: %s\n", assetsDir) + fmt.Fprintf(os.Stderr, "Database: %s\n\n", maskPassword(assetsConn)) + + fmt.Fprintf(os.Stderr, "[1/2] Scanning asset manifests...\n") + + items, err := assetloader.ScanDir(assetsDir) + if err != nil { + return fmt.Errorf("scanning directory: %w", err) + } + + if len(items) == 0 { + fmt.Fprintf(os.Stderr, " No asset manifests found. Nothing to execute.\n\n") + return nil + } + + fmt.Fprintf(os.Stderr, " ✓ Found %d asset entry(ies)\n\n", len(items)) + + fmt.Fprintf(os.Stderr, "[2/2] Executing assets in order (Priority → Sequence → Dir)...\n\n") + + ctx := context.Background() + conn, err := pgsql.Connect(ctx, assetsConn, "assets-execute") + if err != nil { + return fmt.Errorf("connecting to database: %w", err) + } + defer conn.Close(ctx) + + successCount := 0 + var failures []struct { + item assetloader.Item + err error + } + + for _, item := range items { + name := filepath.Base(item.Entry.File) + fmt.Printf("Executing asset: %s (Priority=%d, Sequence=%d, Dir=%s)\n", + name, item.Priority, item.Sequence, item.DirName) + + if err := assetloader.ExecuteItem(ctx, conn, item); err != nil { + if assetsIgnoreErrors { + fmt.Printf("⚠ Error loading %s: %v (continuing due to --ignore-errors)\n", name, err) + failures = append(failures, struct { + item assetloader.Item + err error + }{item, err}) + continue + } + return fmt.Errorf("asset %s (Priority=%d, Sequence=%d): %w", + name, item.Priority, item.Sequence, err) + } + + successCount++ + fmt.Printf("✓ Successfully loaded: %s\n", name) + } + + fmt.Fprintf(os.Stderr, "\n=== Execution Complete ===\n") + fmt.Fprintf(os.Stderr, "Completed at: %s\n", getCurrentTimestamp()) + fmt.Fprintf(os.Stderr, "Total entries: %d\n", len(items)) + fmt.Fprintf(os.Stderr, "Successful: %d\n", successCount) + if len(failures) > 0 { + fmt.Fprintf(os.Stderr, "Failed: %d\n", len(failures)) + fmt.Fprintf(os.Stderr, "\n⚠ Failed Entries Summary (%d failed):\n", len(failures)) + for i, f := range failures { + fmt.Fprintf(os.Stderr, " %d. %s (Priority=%d, Sequence=%d)\n Error: %v\n", + i+1, filepath.Base(f.item.Entry.File), f.item.Priority, f.item.Sequence, f.err) + } + } + fmt.Fprintf(os.Stderr, "\n") + + return nil +} diff --git a/cmd/relspec/root.go b/cmd/relspec/root.go index 32ed5b4..fd92c90 100644 --- a/cmd/relspec/root.go +++ b/cmd/relspec/root.go @@ -64,6 +64,7 @@ func init() { rootCmd.AddCommand(diffCmd) rootCmd.AddCommand(inspectCmd) rootCmd.AddCommand(scriptsCmd) + rootCmd.AddCommand(assetsCmd) rootCmd.AddCommand(templCmd) rootCmd.AddCommand(editCmd) rootCmd.AddCommand(mergeCmd) diff --git a/docs/SCRIPTS_COMMAND.md b/docs/SCRIPTS_COMMAND.md index 7bbcbea..4876228 100644 --- a/docs/SCRIPTS_COMMAND.md +++ b/docs/SCRIPTS_COMMAND.md @@ -85,6 +85,23 @@ migrations/ All files will be found and executed in Priority→Sequence order regardless of directory structure. +## External File Embedding + +Script SQL can embed nearby text or binary files before execution using `-- @embed` directives: + +```sql +-- @embed: path=assets/message.txt var=:message mode=text +-- @embed: path=assets/photo.bin var=:payload mode=base64 +INSERT INTO assets (message, payload) +VALUES (:message, decode(:payload, 'base64')::bytea); +``` + +- `path`: File path resolved relative to the SQL file containing the directive +- `var`: Named placeholder to replace, such as `:message` +- `mode`: `text` embeds an escaped SQL string literal; `base64` embeds a base64 string literal + +The directive comment is removed from the SQL, and every matching placeholder is replaced before the script is listed or executed. + ## Commands ### relspec scripts list diff --git a/pkg/assetloader/embed.go b/pkg/assetloader/embed.go new file mode 100644 index 0000000..d8960af --- /dev/null +++ b/pkg/assetloader/embed.go @@ -0,0 +1,165 @@ +package assetloader + +import ( + "encoding/base64" + "fmt" + "os" + "path/filepath" + "regexp" + "strconv" + "strings" + "unicode/utf8" +) + +const ScriptSourcePathMetadataKey = "source_path" + +var ( + embedDirectivePattern = regexp.MustCompile(`(?m)^\s*--\s*@embed:\s*(.+?)\s*$`) + embedAttrPattern = regexp.MustCompile(`([a-zA-Z_][a-zA-Z0-9_]*)=("[^"]*"|'[^']*'|\S+)`) + embedVarPattern = regexp.MustCompile(`^:[a-zA-Z_][a-zA-Z0-9_]*$`) +) + +// ProcessEmbedDirectives expands SQL comments in the form: +// +// -- @embed: path=... var=:... mode=text|base64 +// +// Paths are resolved relative to sqlPath. Text mode embeds a quoted UTF-8 SQL +// string literal. Base64 mode embeds a quoted base64 literal suitable for +// decode(:var, 'base64'). +func ProcessEmbedDirectives(sqlPath, sql string) (string, error) { + directives := embedDirectivePattern.FindAllStringSubmatch(sql, -1) + if len(directives) == 0 { + return sql, nil + } + + if sqlPath == "" { + return "", fmt.Errorf("sql path is required for embed directives") + } + + result := embedDirectivePattern.ReplaceAllString(sql, "") + for i, directive := range directives { + literal, placeholder, err := embedDirectiveLiteral(sqlPath, directive[1], i+1) + if err != nil { + return "", err + } + if !embedPlaceholderPattern(placeholder).MatchString(result) { + return "", fmt.Errorf("%s embed directive %d: placeholder %s not found", sqlPath, i+1, placeholder) + } + result = replaceEmbedPlaceholder(result, placeholder, literal) + } + + return result, nil +} + +func embedDirectiveLiteral(sqlPath, raw string, directiveNumber int) (string, string, error) { + attrs, err := parseEmbedAttrs(raw) + if err != nil { + return "", "", fmt.Errorf("%s embed directive %d: %w", sqlPath, directiveNumber, err) + } + + pathValue := attrs["path"] + varValue := attrs["var"] + modeValue := attrs["mode"] + if pathValue == "" { + return "", "", fmt.Errorf("%s embed directive %d: missing path", sqlPath, directiveNumber) + } + if !embedVarPattern.MatchString(varValue) { + return "", "", fmt.Errorf("%s embed directive %d: var must be a named placeholder like :asset", sqlPath, directiveNumber) + } + if modeValue != "text" && modeValue != "base64" { + return "", "", fmt.Errorf("%s embed directive %d: mode must be text or base64", sqlPath, directiveNumber) + } + + resolved := filepath.Join(filepath.Dir(sqlPath), filepath.Clean(pathValue)) + data, err := os.ReadFile(resolved) + if err != nil { + return "", "", fmt.Errorf("%s embed directive %d: reading %s: %w", sqlPath, directiveNumber, resolved, err) + } + + switch modeValue { + case "text": + if !utf8.Valid(data) { + return "", "", fmt.Errorf("%s embed directive %d: %s is not valid UTF-8", sqlPath, directiveNumber, resolved) + } + literal, err := sqlStringLiteral(string(data)) + if err != nil { + return "", "", fmt.Errorf("%s embed directive %d: %w", sqlPath, directiveNumber, err) + } + return literal, varValue, nil + case "base64": + literal, err := sqlStringLiteral(base64.StdEncoding.EncodeToString(data)) + if err != nil { + return "", "", fmt.Errorf("%s embed directive %d: %w", sqlPath, directiveNumber, err) + } + return literal, varValue, nil + default: + return "", "", fmt.Errorf("%s embed directive %d: mode must be text or base64", sqlPath, directiveNumber) + } +} + +func parseEmbedAttrs(raw string) (map[string]string, error) { + attrs := map[string]string{} + matches := embedAttrPattern.FindAllStringSubmatchIndex(raw, -1) + if len(matches) == 0 { + return nil, fmt.Errorf("expected path, var, and mode attributes") + } + + lastEnd := 0 + for _, match := range matches { + gap := strings.TrimSpace(raw[lastEnd:match[0]]) + if gap != "" { + return nil, fmt.Errorf("invalid attribute syntax near %q", gap) + } + + key := raw[match[2]:match[3]] + value := raw[match[4]:match[5]] + if _, exists := attrs[key]; exists { + return nil, fmt.Errorf("duplicate attribute %q", key) + } + unquoted, err := unquoteEmbedValue(value) + if err != nil { + return nil, fmt.Errorf("invalid %s value: %w", key, err) + } + attrs[key] = unquoted + lastEnd = match[1] + } + if tail := strings.TrimSpace(raw[lastEnd:]); tail != "" { + return nil, fmt.Errorf("invalid attribute syntax near %q", tail) + } + + for key := range attrs { + if key != "path" && key != "var" && key != "mode" { + return nil, fmt.Errorf("unknown attribute %q", key) + } + } + + return attrs, nil +} + +func unquoteEmbedValue(value string) (string, error) { + if len(value) < 2 { + return value, nil + } + if value[0] == '"' { + return strconv.Unquote(value) + } + if value[0] == '\'' && value[len(value)-1] == '\'' { + return value[1 : len(value)-1], nil + } + return value, nil +} + +func sqlStringLiteral(value string) (string, error) { + if strings.ContainsRune(value, '\x00') { + return "", fmt.Errorf("embedded text contains NUL byte") + } + return "'" + strings.ReplaceAll(value, "'", "''") + "'", nil +} + +func replaceEmbedPlaceholder(sql, placeholder, literal string) string { + return embedPlaceholderPattern(placeholder).ReplaceAllString(sql, "${1}"+literal+"${2}") +} + +func embedPlaceholderPattern(placeholder string) *regexp.Regexp { + return regexp.MustCompile(`(^|[^a-zA-Z0-9_:])` + regexp.QuoteMeta(placeholder) + `([^a-zA-Z0-9_]|$)`) +} diff --git a/pkg/assetloader/embed_test.go b/pkg/assetloader/embed_test.go new file mode 100644 index 0000000..062c2c3 --- /dev/null +++ b/pkg/assetloader/embed_test.go @@ -0,0 +1,143 @@ +package assetloader + +import ( + "encoding/base64" + "os" + "path/filepath" + "strings" + "testing" +) + +func TestProcessEmbedDirectives_TextLiteralEscapesQuotes(t *testing.T) { + dir := t.TempDir() + sqlPath := filepath.Join(dir, "1_001_seed.sql") + if err := os.WriteFile(filepath.Join(dir, "body.txt"), []byte("Line 1\nIt's fine"), 0o644); err != nil { + t.Fatal(err) + } + + got, err := ProcessEmbedDirectives(sqlPath, ` +-- @embed: path=body.txt var=:body mode=text +INSERT INTO notes (body) VALUES (:body); +`) + if err != nil { + t.Fatalf("ProcessEmbedDirectives failed: %v", err) + } + if !strings.Contains(got, "VALUES ('Line 1\nIt''s fine');") { + t.Fatalf("embedded SQL did not contain escaped text literal:\n%s", got) + } + if strings.Contains(got, "VALUES (:body);") { + t.Fatalf("placeholder was not replaced:\n%s", got) + } +} + +func TestProcessEmbedDirectives_Base64Literal(t *testing.T) { + dir := t.TempDir() + sqlPath := filepath.Join(dir, "1_001_seed.sql") + binary := []byte{0x00, 0xff, 0x10, 0x20} + if err := os.WriteFile(filepath.Join(dir, "blob.bin"), binary, 0o644); err != nil { + t.Fatal(err) + } + + got, err := ProcessEmbedDirectives(sqlPath, ` +-- @embed: path=blob.bin var=:payload mode=base64 +INSERT INTO files (payload) VALUES (decode(:payload, 'base64')::bytea); +`) + if err != nil { + t.Fatalf("ProcessEmbedDirectives failed: %v", err) + } + want := "decode('" + base64.StdEncoding.EncodeToString(binary) + "', 'base64')::bytea" + if !strings.Contains(got, want) { + t.Fatalf("embedded SQL did not contain base64 literal %q:\n%s", want, got) + } +} + +func TestProcessEmbedDirectives_RelativeToSQLFile(t *testing.T) { + root := t.TempDir() + sqlDir := filepath.Join(root, "nested", "seed") + if err := os.MkdirAll(filepath.Join(sqlDir, "assets"), 0o755); err != nil { + t.Fatal(err) + } + sqlPath := filepath.Join(sqlDir, "1_001_seed.sql") + if err := os.WriteFile(filepath.Join(sqlDir, "assets", "body.txt"), []byte("relative body"), 0o644); err != nil { + t.Fatal(err) + } + + got, err := ProcessEmbedDirectives(sqlPath, ` +-- @embed: path=assets/body.txt var=:body mode=text +SELECT :body; +`) + if err != nil { + t.Fatalf("ProcessEmbedDirectives failed: %v", err) + } + if !strings.Contains(got, "SELECT 'relative body';") { + t.Fatalf("path was not resolved relative to SQL file:\n%s", got) + } +} + +func TestProcessEmbedDirectives_InvalidDirectiveAndFiles(t *testing.T) { + dir := t.TempDir() + sqlPath := filepath.Join(dir, "1_001_seed.sql") + if err := os.WriteFile(filepath.Join(dir, "body.txt"), []byte("ok"), 0o644); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(dir, "binary.txt"), []byte{0xff, 0xfe}, 0o644); err != nil { + t.Fatal(err) + } + + tests := []struct { + name string + sql string + }{ + { + name: "missing mode", + sql: "-- @embed: path=body.txt var=:body\nSELECT :body;", + }, + { + name: "invalid var", + sql: "-- @embed: path=body.txt var=body mode=text\nSELECT :body;", + }, + { + name: "missing file", + sql: "-- @embed: path=missing.txt var=:body mode=text\nSELECT :body;", + }, + { + name: "invalid utf8 text", + sql: "-- @embed: path=binary.txt var=:body mode=text\nSELECT :body;", + }, + { + name: "placeholder not found", + sql: "-- @embed: path=body.txt var=:body mode=text\nSELECT 1;", + }, + { + name: "unknown attribute", + sql: "-- @embed: path=body.txt var=:body mode=text extra=yes\nSELECT :body;", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if _, err := ProcessEmbedDirectives(sqlPath, tt.sql); err == nil { + t.Fatal("expected error, got nil") + } + }) + } +} + +func TestProcessEmbedDirectives_DoesNotReplacePlaceholderPrefix(t *testing.T) { + dir := t.TempDir() + sqlPath := filepath.Join(dir, "1_001_seed.sql") + if err := os.WriteFile(filepath.Join(dir, "body.txt"), []byte("ok"), 0o644); err != nil { + t.Fatal(err) + } + + got, err := ProcessEmbedDirectives(sqlPath, ` +-- @embed: path=body.txt var=:body mode=text +SELECT :body, :body_extra; +`) + if err != nil { + t.Fatalf("ProcessEmbedDirectives failed: %v", err) + } + if !strings.Contains(got, "SELECT 'ok', :body_extra;") { + t.Fatalf("placeholder boundary was not respected:\n%s", got) + } +} diff --git a/pkg/assetloader/executor.go b/pkg/assetloader/executor.go new file mode 100644 index 0000000..236fdc7 --- /dev/null +++ b/pkg/assetloader/executor.go @@ -0,0 +1,107 @@ +package assetloader + +import ( + "context" + "fmt" + "os" + "path/filepath" + "regexp" + "strings" + + "github.com/jackc/pgx/v5" +) + +// namedPlaceholder matches :identifier patterns (but not ::cast syntax). +var namedPlaceholder = regexp.MustCompile(`:([a-zA-Z_][a-zA-Z0-9_]*)`) + +// pgCastMarker temporarily replaces :: to protect PostgreSQL cast syntax. +const pgCastMarker = "\x00PGCAST\x00" + +// BuildQuery converts a SQL call that uses :name named placeholders into a +// pgx-compatible positional-parameter query ($1, $2, …) and returns the +// corresponding argument slice. +// +// Built-in placeholders: +// - :bytes → fileBytes ([]byte) +// - :filename → filename (string, base name only) +// - :any_key → staticParams["any_key"] (string) +// +// A placeholder that appears more than once maps to the same $N. An unknown +// placeholder (not built-in and not in staticParams) returns an error. +// PostgreSQL cast syntax (::type) is left untouched. +func BuildQuery(call string, fileBytes []byte, filename string, staticParams map[string]string) (string, []any, error) { + // Protect :: casts before running the placeholder regex. + protected := strings.ReplaceAll(call, "::", pgCastMarker) + + paramIndex := map[string]int{} // name → 1-based position + var args []any + var firstErr error + + result := namedPlaceholder.ReplaceAllStringFunc(protected, func(match string) string { + if firstErr != nil { + return match + } + name := match[1:] // strip leading ':' + + // Return existing positional param for repeated placeholders. + if idx, seen := paramIndex[name]; seen { + return fmt.Sprintf("$%d", idx) + } + + // Resolve the placeholder value. + var val any + switch name { + case "bytes": + val = fileBytes + case "filename": + val = filename + default: + if staticParams != nil { + if v, ok := staticParams[name]; ok { + val = v + } + } + if val == nil { + firstErr = fmt.Errorf("unknown placeholder %q in SQL call (not a built-in and not listed in params)", match) + return match + } + } + + idx := len(args) + 1 + paramIndex[name] = idx + args = append(args, val) + return fmt.Sprintf("$%d", idx) + }) + + if firstErr != nil { + return "", nil, firstErr + } + + // Restore :: casts. + result = strings.ReplaceAll(result, pgCastMarker, "::") + + return result, args, nil +} + +// ExecuteItem reads the asset file referenced by item.Entry.File (which is the +// absolute path set by ScanDir) and executes the configured SQL call via conn. +// The file's raw bytes are bound as a []byte parameter — no encoding or escaping. +func ExecuteItem(ctx context.Context, conn *pgx.Conn, item Item) error { + data, err := os.ReadFile(item.Entry.File) + if err != nil { + return fmt.Errorf("reading asset file %s: %w", item.Entry.File, err) + } + + filename := filepath.Base(item.Entry.File) + + sql, args, err := BuildQuery(item.Entry.Call, data, filename, item.Entry.Params) + if err != nil { + return fmt.Errorf("building query for %s: %w", filename, err) + } + + if _, err := conn.Exec(ctx, sql, args...); err != nil { + return fmt.Errorf("executing asset %s: %w", filename, err) + } + + return nil +} diff --git a/pkg/assetloader/loader.go b/pkg/assetloader/loader.go new file mode 100644 index 0000000..7fee0d0 --- /dev/null +++ b/pkg/assetloader/loader.go @@ -0,0 +1,200 @@ +// Package assetloader implements a native Go asset/file loader that binds +// local binary and text files as pgx query parameters during database seeding. +// Files are bound as actual []byte query parameters — never converted to SQL +// text literals — so binary data stays byte-exact and no escaping is needed. +// +// Manifests are small YAML files (assets.yaml) that describe, per file, the +// SQL call to invoke and the named placeholders for :bytes, :filename, and any +// static column values. Manifests live inside directories that follow the same +// {priority}_{sequence}_{name} naming convention used by the sqldir reader, +// so asset-loading steps can be interleaved with SQL scripts in a migrate-apply +// run. +package assetloader + +import ( + "fmt" + "os" + "path/filepath" + "regexp" + "sort" + "strconv" + "strings" + + "gopkg.in/yaml.v3" +) + +// ManifestEntry describes a single file to load from an assets.yaml manifest. +type ManifestEntry struct { + // File is the path to the asset file, relative to the manifest directory. + File string `yaml:"file"` + // Call is the SQL statement to execute. Use :bytes for file content, + // :filename for the base name, and :param_name for static params. + Call string `yaml:"call"` + // Params holds optional static named parameters referenced in Call. + Params map[string]string `yaml:"params,omitempty"` +} + +// Item combines a manifest entry with its ordering metadata and the resolved +// directory where the manifest and asset file reside. +type Item struct { + // Priority and Sequence come from the parent directory's naming pattern. + Priority int + Sequence uint + // DirName is the last path component of the manifest's directory. + DirName string + // Dir is the absolute path to the directory containing assets.yaml and files. + Dir string + // Entry is the parsed manifest entry. + Entry ManifestEntry +} + +// dirPattern matches {priority}_{sequence}_{name} or {priority}-{sequence}-{name} +// directory names, e.g. "1_010_seed_templates" or "2-001-branding". +var dirPattern = regexp.MustCompile(`^(\d+)[_-](\d+)[_-](.+)$`) + +// LoadManifest reads and parses the assets.yaml file in dir, returning +// the ordered list of manifest entries. Returns an error if assets.yaml is +// absent or contains invalid YAML. +func LoadManifest(dir string) ([]ManifestEntry, error) { + manifestPath := filepath.Join(dir, "assets.yaml") + data, err := os.ReadFile(manifestPath) + if err != nil { + return nil, fmt.Errorf("reading %s: %w", manifestPath, err) + } + + var entries []ManifestEntry + if err := yaml.Unmarshal(data, &entries); err != nil { + return nil, fmt.Errorf("parsing %s: %w", manifestPath, err) + } + + return entries, nil +} + +// ScanDir recursively walks baseDir, finds all assets.yaml manifests, resolves +// each file entry (skipping symlinks and path traversal), and returns the +// resulting Items sorted by (Priority, Sequence, DirName). +// +// Each manifest must reside in a directory whose name follows the +// {priority}_{sequence}_{name} pattern. Manifests in directories that do not +// follow this convention are assigned Priority=0, Sequence=0 and sorted last. +func ScanDir(baseDir string) ([]Item, error) { + absBase, err := filepath.Abs(baseDir) + if err != nil { + return nil, fmt.Errorf("resolving base dir: %w", err) + } + + var items []Item + + err = filepath.WalkDir(absBase, func(path string, d os.DirEntry, walkErr error) error { + if walkErr != nil { + return walkErr + } + if d.IsDir() { + return nil + } + if d.Name() != "assets.yaml" { + return nil + } + + manifestDir := filepath.Dir(path) + priority, sequence, dirName := parseDirName(filepath.Base(manifestDir)) + + entries, err := LoadManifest(manifestDir) + if err != nil { + return fmt.Errorf("loading manifest in %s: %w", manifestDir, err) + } + + for _, entry := range entries { + if entry.File == "" || entry.Call == "" { + continue + } + + // Resolve and validate the asset file path. + resolved, skip, err := resolveAssetPath(absBase, manifestDir, entry.File) + if err != nil { + return err + } + if skip { + continue + } + + items = append(items, Item{ + Priority: priority, + Sequence: sequence, + DirName: dirName, + Dir: manifestDir, + Entry: ManifestEntry{ + File: resolved, // absolute path, safe to read + Call: entry.Call, + Params: entry.Params, + }, + }) + } + + return nil + }) + if err != nil { + return nil, err + } + + sort.SliceStable(items, func(i, j int) bool { + if items[i].Priority != items[j].Priority { + return items[i].Priority < items[j].Priority + } + if items[i].Sequence != items[j].Sequence { + return items[i].Sequence < items[j].Sequence + } + return items[i].DirName < items[j].DirName + }) + + return items, nil +} + +// parseDirName extracts (priority, sequence, name) from a directory name that +// follows the {priority}[_-]{sequence}[_-]{name} convention. Returns (0, 0, dir) +// when the name does not match. +func parseDirName(dir string) (priority int, sequence uint, name string) { + m := dirPattern.FindStringSubmatch(dir) + if m == nil { + return 0, 0, dir + } + p, _ := strconv.Atoi(m[1]) + s, _ := strconv.ParseUint(m[2], 10, 64) + return p, uint(s), m[3] +} + +// resolveAssetPath resolves a manifest-relative file path and checks that: +// - it does not escape the base directory (path traversal prevention) +// - none of its path components are symlinks +// +// Returns the absolute path, a skip flag (true when the entry should be silently +// dropped), and any hard error. +func resolveAssetPath(absBase, manifestDir, file string) (absPath string, skip bool, err error) { + // Clean and join before any symlink resolution so we can detect traversal. + joined := filepath.Join(manifestDir, filepath.Clean(file)) + + // Ensure the cleaned path is still inside absBase. + rel, err := filepath.Rel(absBase, joined) + if err != nil || strings.HasPrefix(rel, "..") { + // Path escapes the base directory; skip silently. + return "", true, nil + } + + // Walk each component to detect symlinks. + parts := strings.Split(rel, string(filepath.Separator)) + current := absBase + for _, part := range parts { + current = filepath.Join(current, part) + info, statErr := os.Lstat(current) + if statErr != nil { + // File doesn't exist; skip. + return "", true, nil + } + if info.Mode()&os.ModeSymlink != 0 { + // Symlink in path; skip silently. + return "", true, nil + } + } + + return joined, false, nil +} diff --git a/pkg/assetloader/loader_test.go b/pkg/assetloader/loader_test.go new file mode 100644 index 0000000..fa26818 --- /dev/null +++ b/pkg/assetloader/loader_test.go @@ -0,0 +1,341 @@ +package assetloader_test + +import ( + "os" + "path/filepath" + "testing" + + "git.warky.dev/wdevs/relspecgo/pkg/assetloader" +) + +func TestLoadManifest_ValidList(t *testing.T) { + dir := t.TempDir() + writeFile(t, dir, "assets.yaml", ` +- file: hello.txt + call: INSERT INTO files (name, data) VALUES (:filename, :bytes) +- file: logo.png + call: UPDATE branding SET logo = :bytes WHERE id = 1 +`) + writeFile(t, dir, "hello.txt", "hello world") + writeFile(t, dir, "logo.png", "\x89PNG\r\n\x1a\n") + + m, err := assetloader.LoadManifest(dir) + if err != nil { + t.Fatalf("LoadManifest failed: %v", err) + } + if len(m) != 2 { + t.Fatalf("expected 2 entries, got %d", len(m)) + } + if m[0].File != "hello.txt" { + t.Errorf("entry 0 file: got %q, want %q", m[0].File, "hello.txt") + } + if m[1].File != "logo.png" { + t.Errorf("entry 1 file: got %q, want %q", m[1].File, "logo.png") + } +} + +func TestLoadManifest_WithStaticParams(t *testing.T) { + dir := t.TempDir() + writeFile(t, dir, "assets.yaml", ` +- file: template.md + call: INSERT INTO templates (owner_id, name, data) VALUES (:owner_id, :filename, :bytes) + params: + owner_id: "42" +`) + writeFile(t, dir, "template.md", "# Template") + + m, err := assetloader.LoadManifest(dir) + if err != nil { + t.Fatalf("LoadManifest failed: %v", err) + } + if len(m) != 1 { + t.Fatalf("expected 1 entry, got %d", len(m)) + } + if m[0].Params["owner_id"] != "42" { + t.Errorf("static param owner_id: got %q, want %q", m[0].Params["owner_id"], "42") + } +} + +func TestLoadManifest_MissingFile(t *testing.T) { + dir := t.TempDir() + // No assets.yaml present + _, err := assetloader.LoadManifest(dir) + if err == nil { + t.Fatal("expected error for missing assets.yaml, got nil") + } +} + +func TestLoadManifest_InvalidYAML(t *testing.T) { + dir := t.TempDir() + writeFile(t, dir, "assets.yaml", `{not: [valid yaml`) + _, err := assetloader.LoadManifest(dir) + if err == nil { + t.Fatal("expected error for invalid YAML, got nil") + } +} + +func TestScanDir_FindsManifests(t *testing.T) { + root := t.TempDir() + + // Directory named with priority-sequence pattern + dir1 := filepath.Join(root, "1_010_seed_templates") + if err := os.MkdirAll(dir1, 0o755); err != nil { + t.Fatal(err) + } + writeFile(t, dir1, "assets.yaml", ` +- file: a.txt + call: INSERT INTO t (data) VALUES (:bytes) +`) + writeFile(t, dir1, "a.txt", "aaa") + + dir2 := filepath.Join(root, "2_001_branding") + if err := os.MkdirAll(dir2, 0o755); err != nil { + t.Fatal(err) + } + writeFile(t, dir2, "assets.yaml", ` +- file: logo.png + call: UPDATE branding SET logo = :bytes +`) + writeFile(t, dir2, "logo.png", "PNG") + + items, err := assetloader.ScanDir(root) + if err != nil { + t.Fatalf("ScanDir failed: %v", err) + } + if len(items) != 2 { + t.Fatalf("expected 2 items, got %d", len(items)) + } + // Should be ordered by priority then sequence + if items[0].Priority != 1 || items[0].Sequence != 10 { + t.Errorf("item[0]: got priority=%d seq=%d, want 1,10", items[0].Priority, items[0].Sequence) + } + if items[1].Priority != 2 || items[1].Sequence != 1 { + t.Errorf("item[1]: got priority=%d seq=%d, want 2,1", items[1].Priority, items[1].Sequence) + } +} + +func TestScanDir_OrdersByPriorityThenSequence(t *testing.T) { + root := t.TempDir() + + for _, d := range []string{"2_002_b", "1_001_a", "2_001_c", "1_002_d"} { + dir := filepath.Join(root, d) + if err := os.MkdirAll(dir, 0o755); err != nil { + t.Fatal(err) + } + writeFile(t, dir, "assets.yaml", ` +- file: x.txt + call: SELECT :bytes +`) + writeFile(t, dir, "x.txt", "x") + } + + items, err := assetloader.ScanDir(root) + if err != nil { + t.Fatalf("ScanDir failed: %v", err) + } + if len(items) != 4 { + t.Fatalf("expected 4 items, got %d", len(items)) + } + + type ps struct{ p int; s uint } + want := []ps{{1, 1}, {1, 2}, {2, 1}, {2, 2}} + for i, w := range want { + got := ps{items[i].Priority, items[i].Sequence} + if got != w { + t.Errorf("items[%d]: got {%d,%d}, want {%d,%d}", i, got.p, got.s, w.p, w.s) + } + } +} + +func TestScanDir_SkipsSymlinks(t *testing.T) { + root := t.TempDir() + + dir1 := filepath.Join(root, "1_001_real") + if err := os.MkdirAll(dir1, 0o755); err != nil { + t.Fatal(err) + } + writeFile(t, dir1, "assets.yaml", ` +- file: a.txt + call: SELECT :bytes +`) + writeFile(t, dir1, "a.txt", "real") + + // Symlink to an asset file - should be skipped during file read + realFile := filepath.Join(root, "real.txt") + writeFile(t, root, "real.txt", "symlink target") + symlink := filepath.Join(dir1, "link.txt") + if err := os.Symlink(realFile, symlink); err != nil { + t.Skip("symlinks not supported:", err) + } + + // Add a manifest entry that references the symlink + writeFile(t, dir1, "assets.yaml", ` +- file: a.txt + call: SELECT :bytes +- file: link.txt + call: SELECT :bytes +`) + + items, err := assetloader.ScanDir(root) + if err != nil { + t.Fatalf("ScanDir failed: %v", err) + } + // The symlink entry should be skipped; only a.txt should remain + if len(items) != 1 { + t.Fatalf("expected 1 item after symlink skip, got %d", len(items)) + } + if filepath.Base(items[0].Entry.File) != "a.txt" { + t.Errorf("expected non-symlink entry, got %q", items[0].Entry.File) + } +} + +func TestScanDir_RejectsPathTraversal(t *testing.T) { + root := t.TempDir() + + dir1 := filepath.Join(root, "1_001_evil") + if err := os.MkdirAll(dir1, 0o755); err != nil { + t.Fatal(err) + } + writeFile(t, dir1, "assets.yaml", ` +- file: ../../etc/passwd + call: SELECT :bytes +`) + + items, err := assetloader.ScanDir(root) + if err != nil { + t.Fatalf("ScanDir failed: %v", err) + } + // Path traversal entry should be skipped + if len(items) != 0 { + t.Fatalf("expected 0 items after path traversal rejection, got %d", len(items)) + } +} + +func TestBuildQuery_BasicPlaceholders(t *testing.T) { + sql, args, err := assetloader.BuildQuery( + "INSERT INTO t (name, data) VALUES (:filename, :bytes)", + []byte("hello"), + "hello.txt", + nil, + ) + if err != nil { + t.Fatalf("BuildQuery failed: %v", err) + } + if sql != "INSERT INTO t (name, data) VALUES ($1, $2)" { + t.Errorf("unexpected SQL: %s", sql) + } + if len(args) != 2 { + t.Fatalf("expected 2 args, got %d", len(args)) + } + if string(args[0].(string)) != "hello.txt" { + t.Errorf("args[0]: got %q, want %q", args[0], "hello.txt") + } + if string(args[1].([]byte)) != "hello" { + t.Errorf("args[1]: got %v, want %v", args[1], []byte("hello")) + } +} + +func TestBuildQuery_StaticParams(t *testing.T) { + sql, args, err := assetloader.BuildQuery( + "INSERT INTO t (owner, name, data) VALUES (:owner_id, :filename, :bytes)", + []byte("data"), + "file.bin", + map[string]string{"owner_id": "99"}, + ) + if err != nil { + t.Fatalf("BuildQuery failed: %v", err) + } + if sql != "INSERT INTO t (owner, name, data) VALUES ($1, $2, $3)" { + t.Errorf("unexpected SQL: %s", sql) + } + if len(args) != 3 { + t.Fatalf("expected 3 args, got %d: %v", len(args), args) + } + if args[0].(string) != "99" { + t.Errorf("args[0]: got %q, want %q", args[0], "99") + } +} + +func TestBuildQuery_RepeatedPlaceholder(t *testing.T) { + sql, args, err := assetloader.BuildQuery( + "SELECT length(:bytes), encode(:bytes, 'base64')", + []byte("abc"), + "f.bin", + nil, + ) + if err != nil { + t.Fatalf("BuildQuery failed: %v", err) + } + // :bytes appears twice but maps to same $1 + if sql != "SELECT length($1), encode($1, 'base64')" { + t.Errorf("unexpected SQL: %s", sql) + } + if len(args) != 1 { + t.Fatalf("expected 1 arg, got %d", len(args)) + } +} + +func TestBuildQuery_PostgresCastNotMatched(t *testing.T) { + // ::text should NOT be treated as a placeholder + sql, args, err := assetloader.BuildQuery( + "SELECT :bytes::text, :filename", + []byte("data"), + "f.txt", + nil, + ) + if err != nil { + t.Fatalf("BuildQuery failed: %v", err) + } + if sql != "SELECT $1::text, $2" { + t.Errorf("unexpected SQL: %s", sql) + } + if len(args) != 2 { + t.Fatalf("expected 2 args, got %d", len(args)) + } +} + +func TestBuildQuery_UnknownPlaceholder(t *testing.T) { + _, _, err := assetloader.BuildQuery( + "SELECT :unknown_param", + []byte("data"), + "f.txt", + nil, + ) + if err == nil { + t.Fatal("expected error for unknown placeholder, got nil") + } +} + +func TestBuildQuery_BinaryFileByteExact(t *testing.T) { + // Binary data with null bytes, high bytes - must pass through unchanged + binary := []byte{0x00, 0xFF, 0x80, 0x01, 0xFE} + _, args, err := assetloader.BuildQuery( + "INSERT INTO blobs (data) VALUES (:bytes)", + binary, + "blob.bin", + nil, + ) + if err != nil { + t.Fatalf("BuildQuery failed: %v", err) + } + if len(args) != 1 { + t.Fatalf("expected 1 arg") + } + got := args[0].([]byte) + if len(got) != len(binary) { + t.Fatalf("byte count: got %d, want %d", len(got), len(binary)) + } + for i, b := range binary { + if got[i] != b { + t.Errorf("byte[%d]: got 0x%02x, want 0x%02x", i, got[i], b) + } + } +} + +// writeFile is a test helper that writes content to a file. +func writeFile(t *testing.T, dir, name, content string) { + t.Helper() + if err := os.WriteFile(filepath.Join(dir, name), []byte(content), 0o644); err != nil { + t.Fatalf("writeFile %s: %v", name, err) + } +} diff --git a/pkg/models/models.go b/pkg/models/models.go index 2645f4b..99f29d2 100644 --- a/pkg/models/models.go +++ b/pkg/models/models.go @@ -350,16 +350,17 @@ const ( // Script represents a database migration or initialization script. // Scripts can have dependencies and rollback capabilities. type Script struct { - Name string `json:"name" yaml:"name" xml:"name"` - Description string `json:"description" yaml:"description" xml:"description"` - SQL string `json:"sql" yaml:"sql" xml:"sql"` - Rollback string `json:"rollback,omitempty" yaml:"rollback,omitempty" xml:"rollback,omitempty"` - RunAfter []string `json:"run_after,omitempty" yaml:"run_after,omitempty" xml:"run_after,omitempty"` - Schema string `json:"schema,omitempty" yaml:"schema,omitempty" xml:"schema,omitempty"` - Version string `json:"version,omitempty" yaml:"version,omitempty" xml:"version,omitempty"` - Priority int `json:"priority,omitempty" yaml:"priority,omitempty" xml:"priority,omitempty"` - Sequence uint `json:"sequence,omitempty" yaml:"sequence,omitempty" xml:"sequence,omitempty"` - GUID string `json:"guid" yaml:"guid" xml:"guid"` + Name string `json:"name" yaml:"name" xml:"name"` + Description string `json:"description" yaml:"description" xml:"description"` + SQL string `json:"sql" yaml:"sql" xml:"sql"` + Rollback string `json:"rollback,omitempty" yaml:"rollback,omitempty" xml:"rollback,omitempty"` + RunAfter []string `json:"run_after,omitempty" yaml:"run_after,omitempty" xml:"run_after,omitempty"` + Schema string `json:"schema,omitempty" yaml:"schema,omitempty" xml:"schema,omitempty"` + Version string `json:"version,omitempty" yaml:"version,omitempty" xml:"version,omitempty"` + Priority int `json:"priority,omitempty" yaml:"priority,omitempty" xml:"priority,omitempty"` + Sequence uint `json:"sequence,omitempty" yaml:"sequence,omitempty" xml:"sequence,omitempty"` + GUID string `json:"guid" yaml:"guid" xml:"guid"` + Metadata map[string]any `json:"metadata,omitempty" yaml:"metadata,omitempty" xml:"-"` } // SQLName returns the script name in lowercase for SQL compatibility. @@ -468,6 +469,7 @@ func InitScript(name string) *Script { return &Script{ Name: name, RunAfter: make([]string, 0), + Metadata: make(map[string]any), GUID: uuid.New().String(), } } diff --git a/pkg/readers/sqldir/README.md b/pkg/readers/sqldir/README.md index e5c9248..6464a70 100644 --- a/pkg/readers/sqldir/README.md +++ b/pkg/readers/sqldir/README.md @@ -45,6 +45,21 @@ migrations/ - `1_001_test.txt` - Wrong extension - `readme.md` - Not a SQL file +## External File Embedding + +SQL files can include external files with `-- @embed` directives. File paths are resolved relative to the SQL file being read. + +```sql +-- @embed: path=assets/message.txt var=:message mode=text +-- @embed: path=assets/payload.bin var=:payload mode=base64 +INSERT INTO assets (message, payload) +VALUES (:message, decode(:payload, 'base64')::bytea); +``` + +- `mode=text` reads UTF-8 text and replaces the placeholder with an escaped SQL string literal. +- `mode=base64` reads any bytes and replaces the placeholder with a base64 SQL string literal. +- The placeholder must be named, for example `:message`, and must appear in the SQL body. + ## Usage ### Basic Usage diff --git a/pkg/readers/sqldir/reader.go b/pkg/readers/sqldir/reader.go index 43e6fea..5de0f9e 100644 --- a/pkg/readers/sqldir/reader.go +++ b/pkg/readers/sqldir/reader.go @@ -7,6 +7,7 @@ import ( "regexp" "strconv" + "git.warky.dev/wdevs/relspecgo/pkg/assetloader" "git.warky.dev/wdevs/relspecgo/pkg/models" "git.warky.dev/wdevs/relspecgo/pkg/readers" ) @@ -151,6 +152,10 @@ func (r *Reader) readScripts() ([]*models.Script, error) { if err != nil { return fmt.Errorf("failed to read file %s: %w", path, err) } + sql, err := assetloader.ProcessEmbedDirectives(path, string(content)) + if err != nil { + return err + } // Get relative path from base directory relPath, err := filepath.Rel(r.options.FilePath, path) @@ -161,9 +166,10 @@ func (r *Reader) readScripts() ([]*models.Script, error) { // Create Script model script := models.InitScript(name) script.Description = fmt.Sprintf("SQL script from %s", relPath) - script.SQL = string(content) + script.SQL = sql script.Priority = priority script.Sequence = uint(sequence) + script.Metadata[assetloader.ScriptSourcePathMetadataKey] = path scripts = append(scripts, script) diff --git a/pkg/readers/sqldir/reader_test.go b/pkg/readers/sqldir/reader_test.go index 1aa81ed..50930cf 100644 --- a/pkg/readers/sqldir/reader_test.go +++ b/pkg/readers/sqldir/reader_test.go @@ -1,8 +1,10 @@ package sqldir import ( + "encoding/base64" "os" "path/filepath" + "strings" "testing" "git.warky.dev/wdevs/relspecgo/pkg/readers" @@ -18,12 +20,12 @@ func TestReader_ReadDatabase(t *testing.T) { // Create test SQL files with both underscore and hyphen separators testFiles := map[string]string{ - "1_001_create_users.sql": "CREATE TABLE users (id SERIAL PRIMARY KEY, name TEXT);", - "1_002_create_posts.sql": "CREATE TABLE posts (id SERIAL PRIMARY KEY, user_id INT);", - "2_001_add_indexes.sql": "CREATE INDEX idx_posts_user_id ON posts(user_id);", - "1_003_seed_data.pgsql": "INSERT INTO users (name) VALUES ('Alice'), ('Bob');", + "1_001_create_users.sql": "CREATE TABLE users (id SERIAL PRIMARY KEY, name TEXT);", + "1_002_create_posts.sql": "CREATE TABLE posts (id SERIAL PRIMARY KEY, user_id INT);", + "2_001_add_indexes.sql": "CREATE INDEX idx_posts_user_id ON posts(user_id);", + "1_003_seed_data.pgsql": "INSERT INTO users (name) VALUES ('Alice'), ('Bob');", "10-10-create-newid.pgsql": "CREATE TABLE newid (id SERIAL PRIMARY KEY);", - "2-005-add-column.sql": "ALTER TABLE users ADD COLUMN email TEXT;", + "2-005-add-column.sql": "ALTER TABLE users ADD COLUMN email TEXT;", } for filename, content := range testFiles { @@ -267,10 +269,10 @@ func TestReader_HyphenFormat(t *testing.T) { // Create test files with hyphen separators testFiles := map[string]string{ - "1-001-create-table.sql": "CREATE TABLE test (id INT);", - "1-002-insert-data.pgsql": "INSERT INTO test VALUES (1);", + "1-001-create-table.sql": "CREATE TABLE test (id INT);", + "1-002-insert-data.pgsql": "INSERT INTO test VALUES (1);", "10-10-create-newid.pgsql": "CREATE TABLE newid (id SERIAL);", - "2-005-add-index.sql": "CREATE INDEX idx_test ON test(id);", + "2-005-add-index.sql": "CREATE INDEX idx_test ON test(id);", } for filename, content := range testFiles { @@ -301,10 +303,10 @@ func TestReader_HyphenFormat(t *testing.T) { priority int sequence uint }{ - "create-table": {1, 1}, - "insert-data": {1, 2}, - "add-index": {2, 5}, - "create-newid": {10, 10}, + "create-table": {1, 1}, + "insert-data": {1, 2}, + "add-index": {2, 5}, + "create-newid": {10, 10}, } for _, script := range schema.Scripts { @@ -435,3 +437,61 @@ func TestReader_SkipSymlinks(t *testing.T) { t.Error("Symlink script should have been skipped but was found") } } + +func TestReader_EmbedDirectives(t *testing.T) { + tempDir := t.TempDir() + assetDir := filepath.Join(tempDir, "assets") + if err := os.MkdirAll(assetDir, 0o755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(assetDir, "message.txt"), []byte("Reader's text"), 0o644); err != nil { + t.Fatal(err) + } + binary := []byte{0x00, 0x01, 0xfe, 0xff} + if err := os.WriteFile(filepath.Join(assetDir, "payload.bin"), binary, 0o644); err != nil { + t.Fatal(err) + } + sql := ` +-- @embed: path=assets/message.txt var=:message mode=text +-- @embed: path=assets/payload.bin var=:payload mode=base64 +INSERT INTO assets (message, payload) VALUES (:message, decode(:payload, 'base64')::bytea); +` + if err := os.WriteFile(filepath.Join(tempDir, "1_001_embed.sql"), []byte(sql), 0o644); err != nil { + t.Fatal(err) + } + + reader := NewReader(&readers.ReaderOptions{FilePath: tempDir}) + db, err := reader.ReadDatabase() + if err != nil { + t.Fatalf("ReadDatabase failed: %v", err) + } + if len(db.Schemas[0].Scripts) != 1 { + t.Fatalf("expected 1 script, got %d", len(db.Schemas[0].Scripts)) + } + + got := db.Schemas[0].Scripts[0].SQL + if !strings.Contains(got, "'Reader''s text'") { + t.Fatalf("text asset was not embedded as an escaped SQL literal:\n%s", got) + } + wantBase64 := "decode('" + base64.StdEncoding.EncodeToString(binary) + "', 'base64')::bytea" + if !strings.Contains(got, wantBase64) { + t.Fatalf("binary asset was not embedded as a base64 SQL literal:\n%s", got) + } +} + +func TestReader_EmbedDirectiveErrors(t *testing.T) { + tempDir := t.TempDir() + sql := "-- @embed: path=missing.txt var=:message mode=text\nSELECT :message;" + if err := os.WriteFile(filepath.Join(tempDir, "1_001_embed.sql"), []byte(sql), 0o644); err != nil { + t.Fatal(err) + } + + reader := NewReader(&readers.ReaderOptions{FilePath: tempDir}) + _, err := reader.ReadDatabase() + if err == nil { + t.Fatal("expected embed error, got nil") + } + if !strings.Contains(err.Error(), "missing.txt") { + t.Fatalf("expected missing file in error, got %v", err) + } +} diff --git a/pkg/writers/sqlexec/writer.go b/pkg/writers/sqlexec/writer.go index 267f340..a3fa847 100644 --- a/pkg/writers/sqlexec/writer.go +++ b/pkg/writers/sqlexec/writer.go @@ -7,6 +7,7 @@ import ( "github.com/jackc/pgx/v5" + "git.warky.dev/wdevs/relspecgo/pkg/assetloader" "git.warky.dev/wdevs/relspecgo/pkg/models" "git.warky.dev/wdevs/relspecgo/pkg/pgsql" "git.warky.dev/wdevs/relspecgo/pkg/writers" @@ -138,7 +139,28 @@ func (w *Writer) executeScripts(ctx context.Context, conn *pgx.Conn, scripts []* script.Name, script.Priority, script.Sequence) // Execute the SQL script - _, err := conn.Exec(ctx, script.SQL) + sql, err := processEmbedDirectives(script) + if err != nil { + if ignoreErrors { + fmt.Printf("⚠ Error preparing %s: %v (continuing due to --ignore-errors)\n", script.Name, err) + failedScripts = append(failedScripts, struct { + name string + priority int + sequence uint + err error + }{ + name: script.Name, + priority: script.Priority, + sequence: script.Sequence, + err: err, + }) + continue + } + return fmt.Errorf("script %s (Priority=%d, Sequence=%d): %w", + script.Name, script.Priority, script.Sequence, err) + } + + _, err = conn.Exec(ctx, sql) if err != nil { if ignoreErrors { fmt.Printf("⚠ Error executing %s: %v (continuing due to --ignore-errors)\n", script.Name, err) @@ -179,3 +201,11 @@ func (w *Writer) executeScripts(ctx context.Context, conn *pgx.Conn, scripts []* return nil } + +func processEmbedDirectives(script *models.Script) (string, error) { + sqlPath, _ := script.Metadata[assetloader.ScriptSourcePathMetadataKey].(string) + if sqlPath == "" { + return script.SQL, nil + } + return assetloader.ProcessEmbedDirectives(sqlPath, script.SQL) +} diff --git a/pkg/writers/sqlexec/writer_test.go b/pkg/writers/sqlexec/writer_test.go index 6596f5c..dce74fa 100644 --- a/pkg/writers/sqlexec/writer_test.go +++ b/pkg/writers/sqlexec/writer_test.go @@ -1,8 +1,12 @@ package sqlexec import ( + "os" + "path/filepath" + "strings" "testing" + "git.warky.dev/wdevs/relspecgo/pkg/assetloader" "git.warky.dev/wdevs/relspecgo/pkg/models" "git.warky.dev/wdevs/relspecgo/pkg/writers" ) @@ -216,3 +220,36 @@ func TestWriter_WriteSchema_EmptyScripts(t *testing.T) { // // Verify results // // Cleanup // } + +func TestProcessEmbedDirectives_UsesScriptSourcePath(t *testing.T) { + dir := t.TempDir() + sqlPath := filepath.Join(dir, "1_001_seed.sql") + if err := os.WriteFile(filepath.Join(dir, "body.txt"), []byte("writer text"), 0o644); err != nil { + t.Fatal(err) + } + + script := models.InitScript("seed") + script.SQL = "-- @embed: path=body.txt var=:body mode=text\nSELECT :body;" + script.Metadata[assetloader.ScriptSourcePathMetadataKey] = sqlPath + + got, err := processEmbedDirectives(script) + if err != nil { + t.Fatalf("processEmbedDirectives failed: %v", err) + } + if !strings.Contains(got, "SELECT 'writer text';") { + t.Fatalf("script embed directive was not processed:\n%s", got) + } +} + +func TestProcessEmbedDirectives_NoSourcePathLeavesSQLUnchanged(t *testing.T) { + script := models.InitScript("seed") + script.SQL = "-- @embed: path=missing.txt var=:body mode=text\nSELECT :body;" + + got, err := processEmbedDirectives(script) + if err != nil { + t.Fatalf("processEmbedDirectives failed: %v", err) + } + if got != script.SQL { + t.Fatalf("expected unchanged SQL without source path, got:\n%s", got) + } +} diff --git a/test_data/asset-smoke/1_010_templates/assets.yaml b/test_data/asset-smoke/1_010_templates/assets.yaml new file mode 100644 index 0000000..19fc066 --- /dev/null +++ b/test_data/asset-smoke/1_010_templates/assets.yaml @@ -0,0 +1,6 @@ +- file: hello.md + call: | + INSERT INTO docs (name, content) + VALUES (:filename, :bytes::text) +- file: logo.png + call: UPDATE branding SET logo = :bytes WHERE id = 1 diff --git a/test_data/asset-smoke/1_010_templates/hello.md b/test_data/asset-smoke/1_010_templates/hello.md new file mode 100644 index 0000000..2965834 --- /dev/null +++ b/test_data/asset-smoke/1_010_templates/hello.md @@ -0,0 +1 @@ +# Hello World diff --git a/test_data/asset-smoke/1_010_templates/logo.png b/test_data/asset-smoke/1_010_templates/logo.png new file mode 100644 index 0000000..a8c735a --- /dev/null +++ b/test_data/asset-smoke/1_010_templates/logo.png @@ -0,0 +1 @@ +PNG_PLACEHOLDER diff --git a/test_data/asset-smoke/2_001_branding/assets.yaml b/test_data/asset-smoke/2_001_branding/assets.yaml new file mode 100644 index 0000000..46d99ea --- /dev/null +++ b/test_data/asset-smoke/2_001_branding/assets.yaml @@ -0,0 +1,4 @@ +- file: banner.txt + call: INSERT INTO banners (data) VALUES (:bytes) + params: + owner_id: "1" diff --git a/test_data/asset-smoke/2_001_branding/banner.txt b/test_data/asset-smoke/2_001_branding/banner.txt new file mode 100644 index 0000000..a395181 --- /dev/null +++ b/test_data/asset-smoke/2_001_branding/banner.txt @@ -0,0 +1 @@ +BANNER CONTENT