Files

115 lines
3.5 KiB
Go

package windows
import (
"errors"
"sort"
"strings"
"unicode/utf16"
"unicode/utf8"
)
var (
ErrInvalidEnvironmentKey = errors.New("invalid Windows environment key")
ErrInvalidEnvironmentValue = errors.New("invalid Windows environment value")
ErrDuplicateEnvironmentKey = errors.New("duplicate Windows environment key")
ErrEnvironmentTooLarge = errors.New("Windows environment block is too large")
)
// EnvironmentEntry is an ordered, case-preserving environment assignment.
// Windows compares names case-insensitively, so Folded is never emitted and is
// used only to make duplicate handling deterministic.
type EnvironmentEntry struct {
Key string
Value string
Folded string
}
// MergeEnvironment overlays request values onto the token-derived base
// environment. Shell resolution happens before this function is called; an
// override therefore cannot redirect the configured executable.
func MergeEnvironment(base, overrides map[string]string) ([]EnvironmentEntry, error) {
entries := make(map[string]EnvironmentEntry, len(base)+len(overrides))
add := func(key, value string, override bool) error {
if !validEnvironmentKey(key, override) {
return ErrInvalidEnvironmentKey
}
if strings.IndexByte(value, 0) >= 0 || !utf8.ValidString(value) {
return ErrInvalidEnvironmentValue
}
folded := strings.ToUpper(key)
if _, exists := entries[folded]; exists {
if !override {
return ErrDuplicateEnvironmentKey
}
delete(entries, folded)
}
entries[folded] = EnvironmentEntry{Key: key, Value: value, Folded: folded}
return nil
}
baseKeys := make([]string, 0, len(base))
for key := range base {
baseKeys = append(baseKeys, key)
}
sort.Strings(baseKeys)
for _, key := range baseKeys {
if err := add(key, base[key], false); err != nil {
return nil, err
}
}
overrideKeys := make([]string, 0, len(overrides))
for key := range overrides {
overrideKeys = append(overrideKeys, key)
}
sort.Strings(overrideKeys)
for _, key := range overrideKeys {
if err := add(key, overrides[key], true); err != nil {
return nil, err
}
}
result := make([]EnvironmentEntry, 0, len(entries))
for _, entry := range entries {
result = append(result, entry)
}
sort.Slice(result, func(left, right int) bool {
if result[left].Folded == result[right].Folded {
return result[left].Key < result[right].Key
}
return result[left].Folded < result[right].Folded
})
return result, nil
}
// BuildEnvironmentBlock converts the merged entries to the UTF-16 block
// expected by CreateProcessAsUser. The final two NUL code units are included.
func BuildEnvironmentBlock(base, overrides map[string]string) ([]uint16, error) {
entries, err := MergeEnvironment(base, overrides)
if err != nil {
return nil, err
}
var units []uint16
for _, entry := range entries {
units = append(units, utf16.Encode([]rune(entry.Key+"="+entry.Value))...)
units = append(units, 0)
if len(units) > 32767 {
return nil, ErrEnvironmentTooLarge
}
}
units = append(units, 0)
return units, nil
}
func validEnvironmentKey(key string, override bool) bool {
if key == "" || strings.IndexByte(key, 0) >= 0 || !utf8.ValidString(key) {
return false
}
if override && (strings.ContainsRune(key, '=') || strings.HasPrefix(key, "=")) {
return false
}
if strings.HasPrefix(key, "=") {
// CreateEnvironmentBlock may return drive-current-directory pseudo
// variables such as =C:=C:\\work. They are accepted only as base data.
return !override && strings.Count(key, "=") == 1
}
return !strings.ContainsRune(key, '=')
}