270 lines
7.4 KiB
Go
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
|
|
}
|