269 lines
7.8 KiB
Go
269 lines
7.8 KiB
Go
package workspace
|
|
|
|
import (
|
|
"crypto/sha256"
|
|
"encoding/hex"
|
|
"encoding/json"
|
|
"fmt"
|
|
"path/filepath"
|
|
"sort"
|
|
"strings"
|
|
|
|
"gitea.maximumdirect.net/eric/notarius/internal/core/artifacts"
|
|
"gitea.maximumdirect.net/eric/notarius/internal/framework/pipeline"
|
|
)
|
|
|
|
const digestPrefixLength = 16
|
|
|
|
type Fingerprint struct {
|
|
Name string `json:"name"`
|
|
Value string `json:"value"`
|
|
}
|
|
|
|
type CheckpointIdentityInput struct {
|
|
Pipeline pipeline.ResolvedPipeline
|
|
InputKey string
|
|
RawInputDigest string
|
|
SourceDigest string
|
|
SelectedLanes []string
|
|
RuntimeOverrides []Fingerprint
|
|
References []artifacts.ReferenceProvenance
|
|
ProvenanceFingerprints []Fingerprint
|
|
}
|
|
|
|
type CheckpointIdentity struct {
|
|
Digest string `json:"digest"`
|
|
PipelineID string `json:"pipeline_id"`
|
|
PipelineDigest string `json:"pipeline_digest"`
|
|
InputKey string `json:"input_key"`
|
|
RawInputDigest string `json:"raw_input_digest,omitempty"`
|
|
SourceDigest string `json:"source_digest,omitempty"`
|
|
SelectedLanes []string `json:"selected_lanes,omitempty"`
|
|
RuntimeOverrides []Fingerprint `json:"runtime_overrides,omitempty"`
|
|
ReferenceDigests []Fingerprint `json:"reference_digests,omitempty"`
|
|
ProvenanceFingerprints []Fingerprint `json:"provenance_fingerprints,omitempty"`
|
|
}
|
|
|
|
func NewCheckpointIdentity(input CheckpointIdentityInput) (CheckpointIdentity, error) {
|
|
pipelineID := strings.TrimSpace(input.Pipeline.ID)
|
|
if pipelineID == "" {
|
|
return CheckpointIdentity{}, fmt.Errorf("checkpoint identity pipeline id must not be empty")
|
|
}
|
|
pipelineDigest := strings.TrimSpace(input.Pipeline.Digest)
|
|
if pipelineDigest == "" {
|
|
return CheckpointIdentity{}, fmt.Errorf("checkpoint identity pipeline digest must not be empty")
|
|
}
|
|
inputKey := strings.TrimSpace(input.InputKey)
|
|
if inputKey == "" {
|
|
inputKey = strings.TrimSpace(input.Pipeline.Input.Module)
|
|
}
|
|
if inputKey == "" {
|
|
return CheckpointIdentity{}, fmt.Errorf("checkpoint identity input key must not be empty")
|
|
}
|
|
|
|
rawInputDigest := strings.TrimSpace(input.RawInputDigest)
|
|
sourceDigest := strings.TrimSpace(input.SourceDigest)
|
|
if rawInputDigest == "" && sourceDigest == "" {
|
|
return CheckpointIdentity{}, fmt.Errorf("checkpoint identity raw input digest or source digest must be set")
|
|
}
|
|
|
|
identity := CheckpointIdentity{
|
|
PipelineID: pipelineID,
|
|
PipelineDigest: pipelineDigest,
|
|
InputKey: inputKey,
|
|
RawInputDigest: rawInputDigest,
|
|
SourceDigest: sourceDigest,
|
|
SelectedLanes: normalizedLanes(input.SelectedLanes, input.Pipeline.ArtifactLanes),
|
|
RuntimeOverrides: normalizeFingerprints(input.RuntimeOverrides),
|
|
ReferenceDigests: referenceFingerprints(input.References),
|
|
ProvenanceFingerprints: normalizeFingerprints(input.ProvenanceFingerprints),
|
|
}
|
|
digest, err := identityDigest(identity)
|
|
if err != nil {
|
|
return CheckpointIdentity{}, err
|
|
}
|
|
identity.Digest = digest
|
|
return identity, nil
|
|
}
|
|
|
|
func (s Settings) CheckpointDirectory(identity CheckpointIdentity) (string, error) {
|
|
if !s.ResumeEnabled || strings.TrimSpace(s.CheckpointsRoot) == "" {
|
|
return "", nil
|
|
}
|
|
relative, err := identity.RelativePath()
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
return SafePath(s.CheckpointsRoot, relative)
|
|
}
|
|
|
|
func (i CheckpointIdentity) RelativePath() (string, error) {
|
|
pipelineID, err := safePathComponent(i.PipelineID)
|
|
if err != nil {
|
|
return "", fmt.Errorf("checkpoint identity pipeline id: %w", err)
|
|
}
|
|
inputKey, err := safePathComponent(i.InputKey)
|
|
if err != nil {
|
|
return "", fmt.Errorf("checkpoint identity input key: %w", err)
|
|
}
|
|
sourceDigest := digestPrefix(i.SourceDigest)
|
|
if sourceDigest == "" {
|
|
sourceDigest = digestPrefix(i.RawInputDigest)
|
|
}
|
|
if sourceDigest == "" {
|
|
return "", fmt.Errorf("checkpoint identity source digest prefix must not be empty")
|
|
}
|
|
pipelineDigest := digestPrefix(i.PipelineDigest)
|
|
if pipelineDigest == "" {
|
|
return "", fmt.Errorf("checkpoint identity pipeline digest prefix must not be empty")
|
|
}
|
|
sourceComponent, err := safePathComponent(sourceDigest)
|
|
if err != nil {
|
|
return "", fmt.Errorf("checkpoint identity source digest: %w", err)
|
|
}
|
|
pipelineComponent, err := safePathComponent(pipelineDigest)
|
|
if err != nil {
|
|
return "", fmt.Errorf("checkpoint identity pipeline digest: %w", err)
|
|
}
|
|
return filepath.ToSlash(filepath.Join(pipelineID, inputKey+"-"+sourceComponent, pipelineComponent)), nil
|
|
}
|
|
|
|
func identityDigest(identity CheckpointIdentity) (string, error) {
|
|
payload := identity
|
|
payload.Digest = ""
|
|
data, err := json.Marshal(payload)
|
|
if err != nil {
|
|
return "", fmt.Errorf("marshal checkpoint identity: %w", err)
|
|
}
|
|
sum := sha256.Sum256(data)
|
|
return "sha256:" + hex.EncodeToString(sum[:]), nil
|
|
}
|
|
|
|
func normalizedLanes(selected []string, resolved []pipeline.ResolvedArtifactLane) []string {
|
|
if len(selected) > 0 {
|
|
return normalizeStrings(selected)
|
|
}
|
|
lanes := make([]string, 0, len(resolved))
|
|
for _, lane := range resolved {
|
|
lanes = append(lanes, lane.ID)
|
|
}
|
|
return normalizeStrings(lanes)
|
|
}
|
|
|
|
func normalizeFingerprints(values []Fingerprint) []Fingerprint {
|
|
if len(values) == 0 {
|
|
return nil
|
|
}
|
|
byName := make(map[string]string, len(values))
|
|
for _, value := range values {
|
|
name := strings.TrimSpace(value.Name)
|
|
fingerprint := strings.TrimSpace(value.Value)
|
|
if name == "" || fingerprint == "" {
|
|
continue
|
|
}
|
|
byName[name] = fingerprint
|
|
}
|
|
if len(byName) == 0 {
|
|
return nil
|
|
}
|
|
names := make([]string, 0, len(byName))
|
|
for name := range byName {
|
|
names = append(names, name)
|
|
}
|
|
sort.Strings(names)
|
|
out := make([]Fingerprint, 0, len(names))
|
|
for _, name := range names {
|
|
out = append(out, Fingerprint{Name: name, Value: byName[name]})
|
|
}
|
|
return out
|
|
}
|
|
|
|
func referenceFingerprints(references []artifacts.ReferenceProvenance) []Fingerprint {
|
|
if len(references) == 0 {
|
|
return nil
|
|
}
|
|
values := make([]Fingerprint, 0, len(references))
|
|
for _, reference := range references {
|
|
digest := strings.TrimSpace(reference.Digest)
|
|
if digest == "" {
|
|
continue
|
|
}
|
|
parts := []string{
|
|
strings.TrimSpace(reference.Stage),
|
|
strings.TrimSpace(reference.LaneID),
|
|
strings.TrimSpace(reference.SlotName),
|
|
strings.TrimSpace(reference.OriginURI),
|
|
}
|
|
values = append(values, Fingerprint{
|
|
Name: strings.Join(parts, ":"),
|
|
Value: digest,
|
|
})
|
|
}
|
|
return normalizeFingerprints(values)
|
|
}
|
|
|
|
func normalizeStrings(values []string) []string {
|
|
if len(values) == 0 {
|
|
return nil
|
|
}
|
|
seen := make(map[string]struct{}, len(values))
|
|
for _, value := range values {
|
|
value = strings.TrimSpace(value)
|
|
if value == "" {
|
|
continue
|
|
}
|
|
seen[value] = struct{}{}
|
|
}
|
|
if len(seen) == 0 {
|
|
return nil
|
|
}
|
|
out := make([]string, 0, len(seen))
|
|
for value := range seen {
|
|
out = append(out, value)
|
|
}
|
|
sort.Strings(out)
|
|
return out
|
|
}
|
|
|
|
func digestPrefix(digest string) string {
|
|
digest = strings.TrimSpace(digest)
|
|
if digest == "" {
|
|
return ""
|
|
}
|
|
if idx := strings.Index(digest, ":"); idx >= 0 {
|
|
digest = digest[idx+1:]
|
|
}
|
|
digest = strings.TrimSpace(digest)
|
|
if len(digest) > digestPrefixLength {
|
|
return digest[:digestPrefixLength]
|
|
}
|
|
return digest
|
|
}
|
|
|
|
func safePathComponent(value string) (string, error) {
|
|
value = strings.TrimSpace(value)
|
|
if value == "" {
|
|
return "", fmt.Errorf("must not be empty")
|
|
}
|
|
var b strings.Builder
|
|
for _, r := range value {
|
|
switch {
|
|
case r >= 'a' && r <= 'z':
|
|
b.WriteRune(r)
|
|
case r >= 'A' && r <= 'Z':
|
|
b.WriteRune(r)
|
|
case r >= '0' && r <= '9':
|
|
b.WriteRune(r)
|
|
case r == '-' || r == '_' || r == '.':
|
|
b.WriteRune(r)
|
|
default:
|
|
b.WriteString(fmt.Sprintf("~%x", r))
|
|
}
|
|
}
|
|
encoded := b.String()
|
|
if encoded == "." || encoded == ".." || strings.Contains(encoded, "..") || strings.ContainsAny(encoded, `/\`) {
|
|
return "", fmt.Errorf("%q is not filesystem safe", value)
|
|
}
|
|
return encoded, nil
|
|
}
|