Files
sow-tools/internal/pipeline/extract.go
T

374 lines
10 KiB
Go

package pipeline
import (
"bytes"
"encoding/json"
"errors"
"fmt"
"os"
"path/filepath"
"slices"
"strings"
"gitea.westgate.pw/ShadowsOverWestgate/sow-tools/internal/erf"
"gitea.westgate.pw/ShadowsOverWestgate/sow-tools/internal/gff"
"gitea.westgate.pw/ShadowsOverWestgate/sow-tools/internal/project"
)
type ExtractResult struct {
ModulePath string
HAKPaths []string
DeletedArchivePaths []string
Written int
Overwritten int
Removed int
Skipped int
}
func Extract(p *project.Project, files ...string) (ExtractResult, error) {
var result ExtractResult
var failures []error
desired := map[string]struct{}{}
allowed := make(map[string]bool)
for _, f := range files {
allowed[f] = true
}
shouldExtractMod := len(allowed) == 0 || allowed[p.Config.Module.ResRef+".mod"]
if shouldExtractMod {
modulePath := p.ModuleArchivePath()
input, err := os.Open(modulePath)
if err != nil {
return ExtractResult{}, fmt.Errorf("open module archive: %w", err)
}
defer input.Close()
archive, err := erf.Read(input)
if err != nil {
return ExtractResult{}, fmt.Errorf("read module archive: %w", err)
}
result.ModulePath = modulePath
written, overwritten, skipped, errs := extractArchiveResources(p, archive, desired)
result.Written += written
result.Overwritten += overwritten
result.Skipped += skipped
failures = append(failures, errs...)
}
hakPaths, err := extractHAKPaths(p)
if err != nil {
return result, fmt.Errorf("scan hak archives: %w", err)
}
slices.Sort(hakPaths)
for _, hakPath := range hakPaths {
filename := filepath.Base(hakPath)
if len(allowed) > 0 && !allowed[filename] {
continue
}
input, err := os.Open(hakPath)
if err != nil {
failures = append(failures, fmt.Errorf("open hak archive %s: %w", hakPath, err))
continue
}
hakArchive, err := erf.Read(input)
input.Close()
if err != nil {
failures = append(failures, fmt.Errorf("read hak archive %s: %w", hakPath, err))
continue
}
result.HAKPaths = append(result.HAKPaths, hakPath)
written, overwritten, skipped, errs := extractArchiveResources(p, hakArchive, desired)
result.Written += written
result.Overwritten += overwritten
result.Skipped += skipped
failures = append(failures, errs...)
}
if p.EffectiveConfig().Extract.CleanupStale == nil || *p.EffectiveConfig().Extract.CleanupStale {
removed, errs := cleanupStaleFiles(p, desired)
result.Removed = removed
failures = append(failures, errs...)
}
if len(failures) > 0 {
return result, errors.Join(failures...)
}
if p.Config.Extract.DeleteModuleArchiveAfterSuccess && result.ModulePath != "" {
if err := os.Remove(result.ModulePath); err != nil {
return result, fmt.Errorf("delete consumed module archive %s: %w", result.ModulePath, err)
}
result.DeletedArchivePaths = append(result.DeletedArchivePaths, result.ModulePath)
}
return result, nil
}
func extractArchiveResources(p *project.Project, archive erf.Archive, desired map[string]struct{}) (int, int, int, []error) {
var failures []error
writtenCount := 0
overwrittenCount := 0
skippedCount := 0
ignored := make(map[string]bool)
for _, ext := range p.Config.Extract.IgnoreExtensions {
ext := strings.TrimPrefix(ext, ".")
ignored[ext] = true
}
for _, resource := range archive.Resources {
ext, ok := erf.ExtensionForResourceType(resource.Type)
if !ok {
failures = append(failures, fmt.Errorf("unsupported resource type 0x%04X for %s", resource.Type, resource.Name))
continue
}
if ignored[ext] {
skippedCount++
continue
}
target, data, err := extractedFile(p, resource, ext)
if err != nil {
failures = append(failures, err)
continue
}
desired[target] = struct{}{}
state, err := writeManagedFile(target, data)
if err != nil {
failures = append(failures, err)
continue
}
switch state {
case writeNew:
writtenCount++
case writeOverwritten:
overwrittenCount++
case writeSkipped:
skippedCount++
}
}
return writtenCount, overwrittenCount, skippedCount, failures
}
func extractedFile(p *project.Project, resource erf.Resource, extension string) (string, []byte, error) {
resref := strings.ToLower(resource.Name)
effective := p.EffectiveConfig()
switch extension {
case "nss":
target, err := extractionTarget(p, "paths.source", effective.Paths.Source, p.SourceDir(), filepath.FromSlash(effective.Scripts.SourceDir), resref+".nss")
if err != nil {
return "", nil, err
}
return target, resource.Data, nil
case "utc", "utd", "ute", "uti", "utm", "utp", "uts", "utt", "utw",
"are", "dlg", "fac", "gic", "git", "ifo", "itp", "jrl":
document, err := gff.Read(bytes.NewReader(resource.Data))
if err != nil {
return "", nil, fmt.Errorf("decode gff %s.%s: %w", resource.Name, extension, err)
}
formatted, err := json.MarshalIndent(document, "", " ")
if err != nil {
return "", nil, fmt.Errorf("marshal json %s.%s: %w", resource.Name, extension, err)
}
formatted = append(formatted, '\n')
target, err := extractionTarget(p, "paths.source", effective.Paths.Source, p.SourceDir(), sourceSubdir(extension), resref+"."+extension+".json")
if err != nil {
return "", nil, err
}
return target, formatted, nil
default:
target, err := extractionTarget(p, "paths.assets", effective.Paths.Assets, p.AssetsDir(), extension, resref+"."+extension)
if err != nil {
return "", nil, err
}
return target, resource.Data, nil
}
}
func extractionTarget(p *project.Project, field, configured, root string, parts ...string) (string, error) {
if strings.TrimSpace(configured) == "" {
return "", fmt.Errorf("cannot extract resource: %s is not configured", field)
}
cleanRoot := filepath.Clean(root)
if cleanRoot == "." || cleanRoot == string(filepath.Separator) || cleanRoot == filepath.Clean(p.Root) {
return "", fmt.Errorf("cannot extract resource: %s resolves to unsafe extraction root %s", field, cleanRoot)
}
target := filepath.Join(append([]string{cleanRoot}, parts...)...)
rel, err := filepath.Rel(cleanRoot, target)
if err != nil {
return "", fmt.Errorf("resolve extraction target %s: %w", target, err)
}
if rel == ".." || strings.HasPrefix(rel, ".."+string(filepath.Separator)) || filepath.IsAbs(rel) {
return "", fmt.Errorf("refusing to extract resource outside %s: %s", cleanRoot, target)
}
return target, nil
}
func extractHAKPaths(p *project.Project) ([]string, error) {
switch p.EffectiveConfig().Extract.HAKDiscovery {
case "configured_haks":
paths := make([]string, 0, len(p.Config.HAKs))
for _, hak := range p.Config.HAKs {
path := p.HAKArchivePath(hak.Name)
if _, err := os.Stat(path); err == nil {
paths = append(paths, path)
} else if err != nil && !errors.Is(err, os.ErrNotExist) {
return nil, err
}
}
slices.Sort(paths)
return paths, nil
default:
hakPaths, err := filepath.Glob(filepath.Join(p.BuildDir(), "*.hak"))
if err != nil {
return nil, err
}
slices.Sort(hakPaths)
return hakPaths, nil
}
}
type writeState int
const (
writeSkipped writeState = iota
writeNew
writeOverwritten
)
func writeManagedFile(path string, data []byte) (writeState, error) {
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
return writeSkipped, fmt.Errorf("create parent directory for %s: %w", path, err)
}
existing, err := os.ReadFile(path)
if err == nil {
if bytes.Equal(existing, data) {
return writeSkipped, nil
}
if err := os.WriteFile(path, data, 0o644); err != nil {
return writeSkipped, fmt.Errorf("overwrite %s: %w", path, err)
}
return writeOverwritten, nil
}
if !errors.Is(err, os.ErrNotExist) {
return writeSkipped, fmt.Errorf("check existing file %s: %w", path, err)
}
if err := os.WriteFile(path, data, 0o644); err != nil {
return writeSkipped, fmt.Errorf("write %s: %w", path, err)
}
return writeNew, nil
}
func cleanupStaleFiles(p *project.Project, desired map[string]struct{}) (int, []error) {
candidates := make([]string, 0, len(p.Inventory.SourceFiles)+len(p.Inventory.ScriptFiles)+len(p.Inventory.AssetFiles))
if safeCleanupRoot(p.SourceDir(), p.Root) {
for _, rel := range p.Inventory.SourceFiles {
candidates = append(candidates, filepath.Join(p.SourceDir(), filepath.FromSlash(rel)))
}
for _, rel := range p.Inventory.ScriptFiles {
candidates = append(candidates, filepath.Join(p.SourceDir(), filepath.FromSlash(rel)))
}
}
if safeCleanupRoot(p.AssetsDir(), p.Root) {
for _, rel := range p.Inventory.AssetFiles {
candidates = append(candidates, filepath.Join(p.AssetsDir(), filepath.FromSlash(rel)))
}
}
removed := 0
var failures []error
for _, path := range candidates {
if _, keep := desired[path]; keep {
continue
}
if err := os.Remove(path); err != nil {
if errors.Is(err, os.ErrNotExist) {
continue
}
failures = append(failures, fmt.Errorf("remove stale file %s: %w", path, err))
continue
}
removed++
cleanupEmptyParents(filepath.Dir(path), p.SourceDir(), p.AssetsDir())
}
return removed, failures
}
func safeCleanupRoot(root, projectRoot string) bool {
cleanRoot := filepath.Clean(root)
return cleanRoot != "." && cleanRoot != string(filepath.Separator) && cleanRoot != filepath.Clean(projectRoot)
}
func cleanupEmptyParents(dir string, roots ...string) {
for {
if dir == "." || dir == string(filepath.Separator) {
return
}
stop := false
for _, root := range roots {
if dir == root {
stop = true
break
}
}
if stop {
return
}
if err := os.Remove(dir); err != nil {
return
}
dir = filepath.Dir(dir)
}
}
func sourceSubdir(extension string) string {
switch strings.ToLower(extension) {
case "are":
return "areas"
case "dlg":
return "dialogs"
case "fac":
return "factions"
case "gic":
return "instance"
case "git":
return "instance"
case "ifo":
return "module"
case "itp":
return "palettes"
case "jrl":
return "journal"
case "utc":
return filepath.Join("blueprints", "creatures")
case "utd":
return filepath.Join("blueprints", "doors")
case "ute":
return filepath.Join("blueprints", "encounters")
case "uti":
return filepath.Join("blueprints", "items")
case "utm":
return filepath.Join("blueprints", "merchants")
case "utp":
return filepath.Join("blueprints", "placeables")
case "uts":
return filepath.Join("blueprints", "sounds")
case "utt":
return filepath.Join("blueprints", "triggers")
case "utw":
return filepath.Join("blueprints", "waypoints")
default:
return extension
}
}