Files
ThothII/tools/tht/internal/authprojection/format.go
T

270 lines
7.4 KiB
Go

package authprojection
import (
"bytes"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"fmt"
"io"
"strings"
)
type manifest struct {
Version int `json:"version"`
Generation string `json:"generation"`
Mode string `json:"mode"`
CanonicalRevision string `json:"canonicalRevision"`
Files []manifestFile `json:"files"`
}
type manifestFile struct {
Name string `json:"name"`
Size int `json:"size"`
SHA256 string `json:"sha256"`
}
func NewSnapshot(mode string, auth, users []byte) (Snapshot, error) {
if (mode != "local" && mode != "oidc") || len(auth) > maximumAuthBytes || len(users) > maximumUsersBytes {
return Snapshot{}, ErrIntegrity
}
if mode == "local" && users == nil {
return Snapshot{}, ErrIntegrity
}
if mode == "oidc" && users != nil {
return Snapshot{}, ErrIntegrity
}
authCopy := append([]byte(nil), auth...)
usersCopy := append([]byte(nil), users...)
authDigest := sha256.Sum256(authCopy)
usersDigest := sha256.Sum256(usersCopy)
authSHA256 := hex.EncodeToString(authDigest[:])
usersSHA256 := "-"
if mode == "local" {
usersSHA256 = hex.EncodeToString(usersDigest[:])
}
record := "thothii-auth-projection-v1\nmode=" + mode + "\nauth=" + authSHA256 + "\nusers=" + usersSHA256 + "\n"
generationDigest := sha256.Sum256([]byte(record))
generation := hex.EncodeToString(generationDigest[:])
return Snapshot{
Mode: mode,
Auth: authCopy,
Users: usersCopy,
AuthSHA256: authSHA256,
UsersSHA256: usersSHA256,
Generation: generation,
CanonicalRevision: "sha256:" + generation,
}, nil
}
func decodeSelector(data []byte) (Selector, error) {
if len(data) == 0 || len(data) > maximumSelectorBytes || hasDuplicateObjectKeys(data) {
return Selector{}, ErrIntegrity
}
var value Selector
if err := decodeExact(data, &value); err != nil || !validSelector(value) {
return Selector{}, ErrIntegrity
}
canonical, err := encodeSelector(value)
if err != nil || !bytes.Equal(data, canonical) {
return Selector{}, ErrIntegrity
}
return value, nil
}
func encodeSelector(value Selector) ([]byte, error) {
if !validSelector(value) {
return nil, ErrIntegrity
}
encoded, err := json.Marshal(value)
if err != nil {
return nil, fmt.Errorf("%w: selector encoding", ErrIntegrity)
}
return append(encoded, '\n'), nil
}
func validSelector(value Selector) bool {
if value.Version != SchemaVersion || !validHex(value.Transaction, 32) {
return false
}
switch value.State {
case "blocked":
return value.Generation == "" && len(value.PreviousGenerations) == 0
case "ready":
if !validHex(value.Generation, 64) || len(value.PreviousGenerations) > retainedPredecessors {
return false
}
seen := map[string]struct{}{value.Generation: {}}
for _, generation := range value.PreviousGenerations {
if !validHex(generation, 64) {
return false
}
if _, duplicate := seen[generation]; duplicate {
return false
}
seen[generation] = struct{}{}
}
return true
default:
return false
}
}
func newManifest(snapshot Snapshot) (manifest, error) {
if err := validateSnapshot(snapshot); err != nil {
return manifest{}, err
}
files := []manifestFile{{Name: "auth.yaml", Size: len(snapshot.Auth), SHA256: snapshot.AuthSHA256}}
if snapshot.Mode == "local" {
files = append(files, manifestFile{Name: "users.yaml", Size: len(snapshot.Users), SHA256: snapshot.UsersSHA256})
}
return manifest{Version: SchemaVersion, Generation: snapshot.Generation, Mode: snapshot.Mode, CanonicalRevision: snapshot.CanonicalRevision, Files: files}, nil
}
func decodeManifest(data []byte) (manifest, error) {
if len(data) == 0 || len(data) > maximumManifestBytes || hasDuplicateObjectKeys(data) {
return manifest{}, ErrIntegrity
}
var value manifest
if err := decodeExact(data, &value); err != nil || !validManifest(value) {
return manifest{}, ErrIntegrity
}
canonical, err := encodeManifest(value)
if err != nil || !bytes.Equal(data, canonical) {
return manifest{}, ErrIntegrity
}
return value, nil
}
func encodeManifest(value manifest) ([]byte, error) {
if !validManifest(value) {
return nil, ErrIntegrity
}
encoded, err := json.Marshal(value)
if err != nil {
return nil, fmt.Errorf("%w: manifest encoding", ErrIntegrity)
}
return append(encoded, '\n'), nil
}
func validateManifest(value manifest, snapshot Snapshot) error {
expected, err := newManifest(snapshot)
if err != nil || !manifestEqual(value, expected) {
return ErrIntegrity
}
return nil
}
func manifestEqual(left, right manifest) bool {
if left.Version != right.Version || left.Generation != right.Generation || left.Mode != right.Mode || left.CanonicalRevision != right.CanonicalRevision || len(left.Files) != len(right.Files) {
return false
}
for index := range left.Files {
if left.Files[index] != right.Files[index] {
return false
}
}
return true
}
func validManifest(value manifest) bool {
if value.Version != SchemaVersion || !validHex(value.Generation, 64) || (value.Mode != "local" && value.Mode != "oidc") || value.CanonicalRevision != "sha256:"+value.Generation {
return false
}
want := []string{"auth.yaml"}
if value.Mode == "local" {
want = append(want, "users.yaml")
}
if len(value.Files) != len(want) {
return false
}
for index, file := range value.Files {
maximum := maximumAuthBytes
if file.Name == "users.yaml" {
maximum = maximumUsersBytes
}
if file.Name != want[index] || file.Size < 0 || file.Size > maximum || !validHex(file.SHA256, 64) {
return false
}
}
return true
}
func validateSnapshot(snapshot Snapshot) error {
expected, err := NewSnapshot(snapshot.Mode, snapshot.Auth, snapshot.Users)
if err != nil || expected.AuthSHA256 != snapshot.AuthSHA256 || expected.UsersSHA256 != snapshot.UsersSHA256 || expected.Generation != snapshot.Generation || expected.CanonicalRevision != snapshot.CanonicalRevision {
return ErrIntegrity
}
return nil
}
func decodeExact(data []byte, destination any) error {
decoder := json.NewDecoder(bytes.NewReader(data))
decoder.DisallowUnknownFields()
if err := decoder.Decode(destination); err != nil {
return err
}
if err := decoder.Decode(&struct{}{}); err != io.EOF {
return fmt.Errorf("multiple JSON values")
}
return nil
}
func hasDuplicateObjectKeys(data []byte) bool {
decoder := json.NewDecoder(bytes.NewReader(data))
return scanJSONValue(decoder, nil) != nil || decoder.Decode(&struct{}{}) != io.EOF
}
func scanJSONValue(decoder *json.Decoder, seen map[string]struct{}) error {
token, err := decoder.Token()
if err != nil {
return err
}
delimiter, ok := token.(json.Delim)
if !ok {
return nil
}
switch delimiter {
case '{':
keys := make(map[string]struct{})
for decoder.More() {
keyToken, err := decoder.Token()
if err != nil {
return err
}
key, ok := keyToken.(string)
if !ok {
return fmt.Errorf("object key")
}
if _, exists := keys[key]; exists {
return fmt.Errorf("duplicate object key")
}
keys[key] = struct{}{}
if err := scanJSONValue(decoder, keys); err != nil {
return err
}
}
_, err := decoder.Token()
return err
case '[':
for decoder.More() {
if err := scanJSONValue(decoder, nil); err != nil {
return err
}
}
_, err := decoder.Token()
return err
default:
return fmt.Errorf("unexpected delimiter")
}
}
func validHex(value string, length int) bool {
if len(value) != length || strings.ToLower(value) != value {
return false
}
_, err := hex.DecodeString(value)
return err == nil
}