Files
ThothII/tools/tht/internal/authconfig/commands.go
T

1060 lines
33 KiB
Go

package authconfig
import (
"bufio"
"bytes"
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"io"
"net"
"net/url"
"os"
"os/exec"
"path/filepath"
"strings"
"sync"
"time"
"unicode"
"unicode/utf16"
"unicode/utf8"
"github.com/aritmolab/thothii/tools/tht/internal/compose"
"github.com/aritmolab/thothii/tools/tht/internal/config"
"github.com/aritmolab/thothii/tools/tht/internal/output"
"github.com/aritmolab/thothii/tools/tht/internal/safeio"
"golang.org/x/term"
"gopkg.in/yaml.v3"
)
const maxPasswordFileBytes int64 = 1025
const authCheckTimeout = 45 * time.Second
const interactiveAuthCheckTimeout = 12 * time.Minute
const maxAuthDiagnosticOutputBytes = 64 * 1024
const maxDevicePromptLineBytes = 4 * 1024
const maxDeviceVerificationURIBytes = 2 * 1024
var errCommandRefused = errors.New("authentication command refused")
var writeNewAuthFile = safeio.WriteCanonicalNewFile
var removeAuthFile = safeio.RemoveCanonicalPrivateRegular
// Run implements the host-only authentication operator surface. It accepts password bytes only
// from an echo-free terminal or a bounded private file, and never writes them to either stream.
func Run(ctx context.Context, installation config.Installation, args []string, stdin io.Reader, stdout, stderr io.Writer) int {
return RunWithRunner(ctx, installation, args, stdin, stdout, stderr, compose.NewRunner(""))
}
// RunWithRunner retains the operator grammar while providing the shared shell-free runner seam
// used by the aggregate doctor and command tests.
func RunWithRunner(ctx context.Context, installation config.Installation, args []string, stdin io.Reader, stdout, stderr io.Writer, runner compose.Runner) int {
if len(args) == 0 {
return authFailure(stderr, "auth requires a subcommand")
}
directory := installation.AuthenticationDirectory()
switch args[0] {
case "configure":
if err := configure(directory, args[1:], stdin, stderr); err != nil {
return authFailure(stderr, authMessage(err))
}
return 0
case "status":
if err := status(directory, args[1:], stdout); err != nil {
return authFailure(stderr, authMessage(err))
}
return 0
case "user":
if err := user(directory, args[1:], stdin, stdout, stderr); err != nil {
return authFailure(stderr, authMessage(err))
}
return 0
case "check":
return checkCommand(ctx, installation, args[1:], stdout, stderr, runner)
default:
return authFailure(stderr, "unknown auth subcommand")
}
}
// AuthDiagnostic is the closed JSON contract emitted by the backend diagnostic command.
type AuthDiagnostic struct {
Level string `json:"level"`
Code string `json:"code"`
Message string `json:"message"`
Field *string `json:"field,omitempty"`
}
// AuthDiagnostics is intentionally isomorphic to backend AuthDiagnostics.
type AuthDiagnostics struct {
Ready bool `json:"ready"`
Mode string `json:"mode"`
Checks []AuthDiagnostic `json:"checks"`
}
type authDiagnosticWire struct {
Level string `json:"level"`
Code string `json:"code"`
Message string `json:"message"`
Field json.RawMessage `json:"field"`
}
type authDiagnosticsWire struct {
Ready *bool `json:"ready"`
Mode string `json:"mode"`
Checks []authDiagnosticWire `json:"checks"`
}
func parseCheckArgs(args []string) (jsonMode, interactive bool, err error) {
for _, arg := range args {
switch arg {
case "--json":
if jsonMode {
return false, false, errCommandRefused
}
jsonMode = true
case "--interactive":
if interactive {
return false, false, errCommandRefused
}
interactive = true
default:
return false, false, errCommandRefused
}
}
return jsonMode, interactive, nil
}
func checkCommand(ctx context.Context, installation config.Installation, args []string, stdout, stderr io.Writer, runner compose.Runner) int {
jsonMode, interactive, err := parseCheckArgs(args)
if err != nil {
return authFailure(stderr, authMessage(err))
}
report, err := runCheck(ctx, installation, runner, interactive, false, stderr)
if err != nil {
fmt.Fprintln(stderr, "tht: authentication diagnostics could not be completed")
return 1
}
if jsonMode {
if err := json.NewEncoder(stdout).Encode(report); err != nil {
fmt.Fprintln(stderr, "tht: authentication diagnostic report could not be written")
return 1
}
} else {
status := "failed"
if report.Ready {
status = "passed"
}
fmt.Fprintf(stdout, "authentication: %s\n", status)
for _, check := range report.Checks {
if check.Level == "error" {
fmt.Fprintf(stdout, "%s: %s\n", check.Code, check.Message)
}
}
}
if report.Ready {
return 0
}
return 1
}
var authDiagnosticCodes = map[string]struct{}{
"auth_ready": {}, "auth_config_incomplete": {}, "auth_config_invalid": {}, "auth_session_store_invalid": {},
"local_user_registry_invalid": {}, "local_admin_missing": {}, "oidc_secret_missing": {},
"oidc_discovery_unreachable": {}, "oidc_issuer_mismatch": {}, "oidc_jwks_unreachable": {},
"oidc_group_catalog_unreachable": {}, "oidc_group_catalog_unauthorized": {}, "oidc_mapped_group_missing": {},
"oidc_mapped_group_ambiguous": {}, "oidc_groups_claim_invalid": {}, "oidc_device_flow_unavailable": {},
}
func validAuthDiagnostics(report AuthDiagnostics) bool {
if report.Mode != "local" && report.Mode != "oidc" && report.Mode != "upstream" && report.Mode != "none" && report.Mode != "mock" {
return false
}
if len(report.Checks) == 0 || len(report.Checks) > 129 {
return false
}
seen := make(map[string]struct{}, len(report.Checks))
hasError := false
for _, check := range report.Checks {
if (check.Level != "error" && check.Level != "info") || !safeDiagnosticText(check.Message) {
return false
}
if _, ok := authDiagnosticCodes[check.Code]; !ok {
return false
}
if check.Field != nil && (!safeDiagnosticText(*check.Field) ||
check.Code != "oidc_mapped_group_missing" && check.Code != "oidc_mapped_group_ambiguous") {
return false
}
field := ""
if check.Field != nil {
field = *check.Field
}
key := check.Code + "\x00" + field
if _, duplicate := seen[key]; duplicate {
return false
}
seen[key] = struct{}{}
hasError = hasError || check.Level == "error"
}
if report.Ready {
return len(report.Checks) == 1 && report.Checks[0].Level == "info" &&
report.Checks[0].Code == "auth_ready" && report.Checks[0].Field == nil
}
if !hasError {
return false
}
for _, check := range report.Checks {
if check.Code == "auth_ready" {
return false
}
}
return true
}
func safeDiagnosticText(value string) bool {
if value == "" || utf16Length(value) > 512 || strings.TrimSpace(value) != value {
return false
}
for _, character := range value {
if unicode.IsControl(character) {
return false
}
}
return true
}
func utf16Length(value string) int {
length := 0
for _, character := range value {
length += utf16.RuneLen(character)
}
return length
}
func decodeAuthDiagnostics(value string) (AuthDiagnostics, error) {
if !utf8.ValidString(value) {
return AuthDiagnostics{}, errors.New("authentication diagnostic report is invalid")
}
decoder := json.NewDecoder(strings.NewReader(value))
decoder.DisallowUnknownFields()
var wire authDiagnosticsWire
if err := decoder.Decode(&wire); err != nil || decoder.Decode(&struct{}{}) != io.EOF {
return AuthDiagnostics{}, errors.New("authentication diagnostic report is invalid")
}
if wire.Ready == nil {
return AuthDiagnostics{}, errors.New("authentication diagnostic report is invalid")
}
report := AuthDiagnostics{Ready: *wire.Ready, Mode: wire.Mode, Checks: make([]AuthDiagnostic, 0, len(wire.Checks))}
for _, item := range wire.Checks {
var field *string
if item.Field != nil {
var decoded string
if string(item.Field) == "null" || json.Unmarshal(item.Field, &decoded) != nil {
return AuthDiagnostics{}, errors.New("authentication diagnostic report is invalid")
}
field = &decoded
}
report.Checks = append(report.Checks, AuthDiagnostic{
Level: item.Level, Code: item.Code, Message: item.Message, Field: field,
})
}
if !validAuthDiagnostics(report) {
return AuthDiagnostics{}, errors.New("authentication diagnostic report is invalid")
}
return report, nil
}
func authenticationSecretValues(installation config.Installation) ([]string, error) {
files, err := installation.SecretFiles()
if err != nil {
return nil, errors.New("authentication diagnostic secret corpus is unavailable")
}
values, err := output.SecretValuesFromFiles(files)
if err != nil {
return nil, errors.New("authentication diagnostic secret corpus is unavailable")
}
return values, nil
}
func sanitizeAuthDiagnostics(report AuthDiagnostics, secrets []string) AuthDiagnostics {
for index := range report.Checks {
report.Checks[index].Message = output.Sanitize(report.Checks[index].Message, secrets)
if report.Checks[index].Field != nil {
field := output.Sanitize(*report.Checks[index].Field, secrets)
report.Checks[index].Field = &field
}
}
return report
}
type devicePromptStream struct {
mu sync.Mutex
pending []byte
discarding bool
emitted bool
secrets []string
output io.Writer
writeErr error
}
func newDevicePromptStream(secrets []string, outputWriter io.Writer) *devicePromptStream {
return &devicePromptStream{
pending: make([]byte, 0, 256), secrets: append([]string(nil), secrets...), output: outputWriter,
}
}
func (stream *devicePromptStream) Observe(chunk []byte) {
stream.mu.Lock()
defer stream.mu.Unlock()
for _, value := range chunk {
if value == '\n' {
if !stream.discarding {
stream.emitLineLocked(string(stream.pending))
}
stream.pending = stream.pending[:0]
stream.discarding = false
continue
}
if stream.discarding {
continue
}
if len(stream.pending) >= maxDevicePromptLineBytes {
stream.pending = stream.pending[:0]
stream.discarding = true
continue
}
stream.pending = append(stream.pending, value)
}
}
func (stream *devicePromptStream) Finish() error {
stream.mu.Lock()
defer stream.mu.Unlock()
if !stream.discarding && len(stream.pending) > 0 {
stream.emitLineLocked(string(stream.pending))
}
stream.pending = nil
stream.discarding = false
return stream.writeErr
}
func (stream *devicePromptStream) emitLineLocked(line string) {
if stream.emitted || stream.output == nil {
return
}
prompt := parseDevicePromptLine(output.Sanitize(line, stream.secrets))
if prompt == "" {
return
}
stream.emitted = true
if _, err := fmt.Fprintln(stream.output, prompt); err != nil {
stream.writeErr = errors.New("authentication device prompt could not be written")
}
}
func parseDevicePromptLine(line string) string {
const prefix = "Open "
const separator = " and enter code "
if len(line) == 0 || len(line) > maxDevicePromptLineBytes || !utf8.ValidString(line) ||
strings.TrimSpace(line) != line || !strings.HasPrefix(line, prefix) {
return ""
}
for _, character := range line {
if unicode.IsControl(character) {
return ""
}
}
payload := strings.TrimPrefix(line, prefix)
if strings.Count(payload, separator) != 1 {
return ""
}
verificationURI, userCode, found := strings.Cut(payload, separator)
if !found || !validDeviceVerificationURI(verificationURI) || !validDeviceUserCode(userCode) {
return ""
}
return prefix + verificationURI + separator + userCode
}
func validDeviceVerificationURI(value string) bool {
if len(value) == 0 || len(value) > maxDeviceVerificationURIBytes || !utf8.ValidString(value) {
return false
}
parsed, err := url.Parse(value)
return err == nil && parsed.IsAbs() && parsed.Scheme == "https" && parsed.Host != "" &&
parsed.Hostname() != "" && parsed.User == nil && parsed.Opaque == "" && parsed.Fragment == "" &&
!strings.ContainsAny(value, " \t\r\n")
}
func validDeviceUserCode(value string) bool {
if len(value) < 4 || len(value) > 256 {
return false
}
for _, character := range value {
letter := (character >= 'A' && character <= 'Z') || (character >= 'a' && character <= 'z')
digit := character >= '0' && character <= '9'
punctuation := character == '-' || character == '.' || character == '_' || character == '~'
if !letter && !digit && !punctuation {
return false
}
}
return true
}
// Check invokes the exact backend diagnostic command. Normal host checks always use one-shot
// Compose execution; aggregate doctor passes useRunningCore=true only after confirming core runs.
func Check(ctx context.Context, installation config.Installation, runner compose.Runner, interactive, useRunningCore bool) (AuthDiagnostics, error) {
return runCheck(ctx, installation, runner, interactive, useRunningCore, nil)
}
func runCheck(ctx context.Context, installation config.Installation, runner compose.Runner, interactive, useRunningCore bool, promptOutput io.Writer) (AuthDiagnostics, error) {
if runner == nil {
return AuthDiagnostics{}, errors.New("authentication diagnostic runner is unavailable")
}
secrets, err := authenticationSecretValues(installation)
if err != nil {
return AuthDiagnostics{}, err
}
timeout := authCheckTimeout
if interactive {
timeout = interactiveAuthCheckTimeout
}
bounded, cancel := context.WithTimeout(ctx, timeout)
defer cancel()
command := []string{"exec", "-T", "core", "node", "dist/auth/diagnostic-command.js", "--json"}
if !useRunningCore {
containerName, err := compose.NewOneShotContainerName("thothii-auth-check")
if err != nil {
return AuthDiagnostics{}, errors.New("authentication diagnostic command failed")
}
command = []string{"run", "--rm", "--no-deps", "--no-TTY", "--name", containerName, "core", "node", "dist/auth/diagnostic-command.js", "--json"}
}
if interactive {
command = append(command, "--interactive")
}
var promptStream *devicePromptStream
var stderrObserver func([]byte)
if interactive && promptOutput != nil {
promptStream = newDevicePromptStream(secrets, promptOutput)
stderrObserver = promptStream.Observe
}
result, err := compose.RunBoundedStreaming(runner, bounded, installation.ComposeArgs(command...), nil, compose.CaptureLimits{
StdoutBytes: maxAuthDiagnosticOutputBytes,
StderrBytes: maxAuthDiagnosticOutputBytes,
}, stderrObserver)
if promptStream != nil {
if promptErr := promptStream.Finish(); promptErr != nil {
return AuthDiagnostics{}, promptErr
}
}
var exitError *exec.ExitError
unsafeLifecycle := errors.Is(err, compose.ErrOutputLimit) || errors.Is(err, compose.ErrProcessReap) ||
errors.Is(err, compose.ErrContainerCleanup) || errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded)
validProcessOutcome := !unsafeLifecycle && (result.ExitCode == 0 && err == nil ||
result.ExitCode == 1 && (err == nil || errors.As(err, &exitError)))
if !validProcessOutcome {
return AuthDiagnostics{}, errors.New("authentication diagnostic command failed")
}
report, decodeErr := decodeAuthDiagnostics(result.Stdout)
if decodeErr != nil {
return AuthDiagnostics{}, errors.New("authentication diagnostic report is invalid")
}
if report.Ready != (result.ExitCode == 0) {
return AuthDiagnostics{}, errors.New("authentication diagnostic command failed")
}
safe := sanitizeAuthDiagnostics(report, secrets)
if !validAuthDiagnostics(safe) {
return AuthDiagnostics{}, errors.New("authentication diagnostic report is invalid")
}
return safe, nil
}
func authFailure(stderr io.Writer, message string) int {
fmt.Fprintf(stderr, "tht: %s\n", message)
return 2
}
func authMessage(err error) string {
switch {
case errors.Is(err, errPasswordFileRequired):
return "--password-file is required when standard input is not an interactive terminal"
case errors.Is(err, errExistingConfiguration):
return "authentication is already configured; refusing to overwrite it"
case errors.Is(err, errOIDCUsers):
return "local user commands are unavailable while OIDC authentication is configured"
case errors.Is(err, errLastAdministrator):
return "refusing to remove or disable the last enabled administrator"
case errors.Is(err, errAmbiguousGroups):
return "OIDC user and administrator groups must be different"
case errors.Is(err, errCommandRefused):
return "authentication command refused"
default:
return "authentication configuration is unavailable or invalid"
}
}
var (
errPasswordFileRequired = errors.New("password file required")
errExistingConfiguration = errors.New("existing configuration")
errOIDCUsers = errors.New("OIDC user commands")
errLastAdministrator = errors.New("last administrator")
errAmbiguousGroups = errors.New("ambiguous groups")
)
type configureRequest struct {
mode, publicURL, adminUser, adminDisplayName, passwordFile string
issuer, clientID, authentikBaseURL, userGroup, adminGroup string
}
func configure(directory string, args []string, stdin io.Reader, stderr io.Writer) error {
request, err := parseConfigure(args)
if err != nil {
return errCommandRefused
}
if err := safeio.EnsurePrivateDirectory(directory); err != nil {
return errCommandRefused
}
if _, _, err := Load(directory); err == nil {
return errExistingConfiguration
}
if err := recoverOrRefuseBootstrapState(directory); err != nil {
return err
}
if request.mode == "oidc" {
if request.userGroup == request.adminGroup {
return errAmbiguousGroups
}
return writeInitialOIDC(directory, request)
}
if request.adminUser == "" {
if err := promptLocalConfigure(&request, stdin, stderr); err != nil {
return err
}
}
password, err := readPassword(request.passwordFile, stdin, stderr)
if err != nil {
return err
}
defer zero(password)
hash, err := HashPassword(password, nil)
if err != nil {
return errCommandRefused
}
initialUser, err := newUser(request.adminUser, request.adminDisplayName, hash, []Role{RoleAdmin})
if err != nil {
return errCommandRefused
}
return writeInitialLocal(directory, request.publicURL, initialUser)
}
func parseConfigure(args []string) (configureRequest, error) {
request := configureRequest{}
values := map[string]*string{
"--mode": &request.mode, "--public-url": &request.publicURL, "--admin-user": &request.adminUser,
"--admin-display-name": &request.adminDisplayName, "--password-file": &request.passwordFile,
"--issuer": &request.issuer, "--client-id": &request.clientID, "--authentik-base-url": &request.authentikBaseURL,
"--user-group": &request.userGroup, "--admin-group": &request.adminGroup,
}
seen := make(map[string]bool, len(values))
for len(args) > 0 {
option := args[0]
args = args[1:]
target, exists := values[option]
if !exists || seen[option] || len(args) == 0 || args[0] == "" {
return configureRequest{}, errCommandRefused
}
*target, args = args[0], args[1:]
seen[option] = true
}
if request.mode != "local" && request.mode != "oidc" || !validPublicURL(request.publicURL) {
return configureRequest{}, errCommandRefused
}
if request.mode == "local" {
if seen["--issuer"] || seen["--client-id"] || seen["--authentik-base-url"] || seen["--user-group"] || seen["--admin-group"] {
return configureRequest{}, errCommandRefused
}
if request.adminUser != "" && request.passwordFile == "" { // TTY supplies the password interactively.
return request, nil
}
return request, nil
}
if seen["--admin-user"] || seen["--admin-display-name"] || seen["--password-file"] || request.issuer == "" || request.clientID == "" || request.authentikBaseURL == "" || request.userGroup == "" || request.adminGroup == "" || !validOIDCURL(request.issuer) || !validHTTPSURL(request.authentikBaseURL) || (!isHTTPS(request.publicURL) && !isLoopbackURL(request.publicURL)) {
return configureRequest{}, errCommandRefused
}
return request, nil
}
func promptLocalConfigure(request *configureRequest, stdin io.Reader, stderr io.Writer) error {
file, ok := stdin.(*os.File)
if !ok || !term.IsTerminal(int(file.Fd())) {
return errPasswordFileRequired
}
reader := bufio.NewReader(file)
username, err := promptLine(reader, stderr, "Initial administrator username: ")
if err != nil || username == "" {
return errCommandRefused
}
displayName, err := promptLine(reader, stderr, "Initial administrator display name (optional): ")
if err != nil {
return errCommandRefused
}
request.adminUser, request.adminDisplayName = username, displayName
return nil
}
func promptLine(reader *bufio.Reader, stderr io.Writer, prompt string) (string, error) {
fmt.Fprint(stderr, prompt)
line, err := reader.ReadString('\n')
if err != nil && !errors.Is(err, io.EOF) {
return "", err
}
return strings.TrimSpace(line), nil
}
func writeInitialLocal(directory, publicURL string, user User) error {
registryBytes, err := yaml.Marshal(Registry{Version: 1, Users: []User{user}})
if err != nil {
return errCommandRefused
}
configBytes, err := yaml.Marshal(Config{Version: 1, Mode: "local", PublicURL: publicURL, Session: defaultSession(), Local: LocalConfig{UsersFile: usersFileName}})
if err != nil {
return errCommandRefused
}
if err := writeNewAuthFile(filepath.Join(directory, usersFileName), append(registryBytes, '\n'), 0o600); err != nil {
return errCommandRefused
}
if err := writeNewAuthFile(filepath.Join(directory, authFileName), append(configBytes, '\n'), 0o600); err != nil {
if removeErr := removeAuthFile(filepath.Join(directory, usersFileName)); removeErr != nil {
return errCommandRefused
}
return errCommandRefused
}
return nil
}
func writeInitialOIDC(directory string, request configureRequest) error {
configuration := Config{
Version: 1, Mode: "oidc", PublicURL: request.publicURL, Session: defaultSession(),
OIDC: OIDCConfig{Issuer: request.issuer, ClientID: request.clientID, ClientSecretRef: "THT_OIDC_CLIENT_SECRET", Scopes: []string{"openid", "profile", "email"}, GroupsClaim: "groups"},
GroupCatalog: GroupCatalogConfig{Driver: "authentik", BaseURL: request.authentikBaseURL, APITokenRef: "THT_AUTHENTIK_API_TOKEN"},
Authorization: AuthorizationConfig{GroupRoles: map[string][]Role{request.userGroup: {RoleUser}, request.adminGroup: {RoleAdmin}}},
}
contents, err := yaml.Marshal(configuration)
if err != nil {
return errCommandRefused
}
if err := writeNewAuthFile(filepath.Join(directory, authFileName), append(contents, '\n'), 0o600); err != nil {
return errCommandRefused
}
return nil
}
func defaultSession() SessionConfig {
return SessionConfig{RegularTTLSeconds: 43200, RegularIdleSeconds: 7200, RememberTTLSeconds: 2592000, RememberIdleSeconds: 604800, OIDCTTLSeconds: 28800}
}
func readPassword(path string, stdin io.Reader, stderr io.Writer) ([]byte, error) {
if path != "" {
if err := safeio.ValidatePrivateRegular(path); err != nil {
return nil, errCommandRefused
}
password, err := safeio.ReadCanonicalRegular(path, maxPasswordFileBytes)
if err != nil || len(password) > passwordMaxBytes+1 {
return nil, errCommandRefused
}
password = trimOneNewline(password)
return password, validatePasswordInput(password)
}
file, ok := stdin.(*os.File)
if !ok || !term.IsTerminal(int(file.Fd())) {
return nil, errPasswordFileRequired
}
fmt.Fprint(stderr, "Password: ")
password, err := term.ReadPassword(int(file.Fd()))
fmt.Fprintln(stderr)
if err != nil {
return nil, errCommandRefused
}
fmt.Fprint(stderr, "Confirm password: ")
confirmation, err := term.ReadPassword(int(file.Fd()))
fmt.Fprintln(stderr)
if err != nil {
zero(password)
return nil, errCommandRefused
}
matched := string(password) == string(confirmation)
zero(confirmation)
if !matched {
zero(password)
return nil, errCommandRefused
}
return password, validatePasswordInput(password)
}
func trimOneNewline(value []byte) []byte {
if bytes.HasSuffix(value, []byte("\r\n")) {
return value[:len(value)-2]
}
if bytes.HasSuffix(value, []byte("\n")) {
return value[:len(value)-1]
}
return value
}
func validatePasswordInput(password []byte) error {
if !validPassword(password) {
return errCommandRefused
}
return nil
}
func zero(value []byte) {
for index := range value {
value[index] = 0
}
}
type statusResult struct {
Mode string `json:"mode"`
PublicURL string `json:"publicUrl"`
UserCounts map[string]int `json:"userCounts"`
Revision string `json:"configRevision"`
}
func status(directory string, args []string, stdout io.Writer) error {
jsonMode := len(args) == 1 && args[0] == "--json"
if len(args) != 0 && !jsonMode {
return errCommandRefused
}
configuration, registry, err := Load(directory)
if err != nil {
return errCommandRefused
}
counts := map[string]int{"user": 0, "admin": 0}
for _, user := range registry.Users {
for _, role := range user.Roles {
counts[string(role)]++
}
}
revision, err := configRevision(directory, configuration.Mode)
if err != nil {
return errCommandRefused
}
result := statusResult{Mode: configuration.Mode, PublicURL: configuration.PublicURL, UserCounts: counts, Revision: revision}
if jsonMode {
encoder := json.NewEncoder(stdout)
encoder.SetEscapeHTML(false)
return encoder.Encode(result)
}
_, err = fmt.Fprintf(stdout, "Authentication mode: %s\nPublic URL: %s\nUsers: user=%d admin=%d\nConfig revision: %s\n", result.Mode, result.PublicURL, counts["user"], counts["admin"], result.Revision)
return err
}
func configRevision(directory, mode string) (string, error) {
auth, err := safeio.ReadCanonicalRegular(filepath.Join(directory, authFileName), maxYAMLBytes)
if err != nil {
return "", err
}
sum := sha256.New()
_, _ = sum.Write(auth)
if mode == "local" {
users, err := safeio.ReadCanonicalRegular(filepath.Join(directory, usersFileName), maxYAMLBytes)
if err != nil {
return "", err
}
_, _ = sum.Write(users)
}
return "sha256:" + hex.EncodeToString(sum.Sum(nil)), nil
}
func user(directory string, args []string, stdin io.Reader, stdout, stderr io.Writer) error {
if len(args) == 0 {
return errCommandRefused
}
configuration, _, err := Load(directory)
if err != nil {
return errCommandRefused
}
if configuration.Mode != "local" {
return errOIDCUsers
}
switch args[0] {
case "list":
return listUsers(directory, args[1:], stdout)
case "add":
return addUser(directory, args[1:], stdin, stderr)
case "set-password":
return setPassword(directory, args[1:], stdin, stderr)
case "enable", "disable":
return setEnabled(directory, args[0] == "enable", args[1:])
case "grant", "revoke":
return changeRole(directory, args[0] == "grant", args[1:])
case "logout-all":
return logoutAll(directory, args[1:])
default:
return errCommandRefused
}
}
func listUsers(directory string, args []string, stdout io.Writer) error {
jsonMode := len(args) == 1 && args[0] == "--json"
if len(args) != 0 && !jsonMode {
return errCommandRefused
}
_, registry, err := Load(directory)
if err != nil {
return errCommandRefused
}
if jsonMode {
return json.NewEncoder(stdout).Encode(struct {
Users []User `json:"users"`
}{Users: registry.Users})
}
for _, item := range registry.Users {
if _, err := fmt.Fprintf(stdout, "%s\t%s\tenabled=%t\troles=%s\n", item.Username, item.DisplayName, item.Enabled, strings.Join(rolesToStrings(item.Roles), ",")); err != nil {
return err
}
}
return nil
}
func addUser(directory string, args []string, stdin io.Reader, stderr io.Writer) error {
if len(args) == 0 {
return errCommandRefused
}
username, request, err := parseUserRolePassword(args)
if err != nil {
return errCommandRefused
}
password, err := readPassword(request.passwordFile, stdin, stderr)
if err != nil {
return err
}
defer zero(password)
hash, err := HashPassword(password, nil)
if err != nil {
return errCommandRefused
}
return MutateUsers(directory, func(registry *Registry) error {
if registry.FindByUsername(username) != nil {
return errCommandRefused
}
item, err := newUser(username, request.displayName, hash, []Role{request.role})
if err != nil {
return errCommandRefused
}
registry.Users = append(registry.Users, item)
return nil
})
}
type userRequest struct {
role Role
displayName, passwordFile string
}
func parseUserRolePassword(args []string) (string, userRequest, error) {
username := args[0]
args = args[1:]
request := userRequest{}
seen := make(map[string]bool, 3)
for len(args) > 0 {
option := args[0]
args = args[1:]
if len(args) == 0 || args[0] == "" || seen[option] {
return "", userRequest{}, errCommandRefused
}
value := args[0]
args = args[1:]
switch option {
case "--role":
request.role = Role(value)
case "--display-name":
request.displayName = value
case "--password-file":
request.passwordFile = value
default:
return "", userRequest{}, errCommandRefused
}
seen[option] = true
}
if !usernamePattern.MatchString(username) || (request.role != RoleUser && request.role != RoleAdmin) {
return "", userRequest{}, errCommandRefused
}
return username, request, nil
}
func setPassword(directory string, args []string, stdin io.Reader, stderr io.Writer) error {
if len(args) < 1 || !usernamePattern.MatchString(args[0]) {
return errCommandRefused
}
username := args[0]
passwordFile, err := parsePasswordFile(args[1:])
if err != nil {
return errCommandRefused
}
password, err := readPassword(passwordFile, stdin, stderr)
if err != nil {
return err
}
defer zero(password)
hash, err := HashPassword(password, nil)
if err != nil {
return errCommandRefused
}
return MutateUsers(directory, func(registry *Registry) error {
item := registry.FindByUsername(username)
if item == nil {
return errCommandRefused
}
item.PasswordHash = hash
return nil
})
}
func parsePasswordFile(args []string) (string, error) {
if len(args) == 0 {
return "", nil
}
if len(args) != 2 || args[0] != "--password-file" || args[1] == "" {
return "", errCommandRefused
}
return args[1], nil
}
func recoverOrRefuseBootstrapState(directory string) error {
authPath := filepath.Join(directory, authFileName)
usersPath := filepath.Join(directory, usersFileName)
_, authErr := os.Lstat(authPath)
_, usersErr := os.Lstat(usersPath)
if !errors.Is(authErr, os.ErrNotExist) {
return errCommandRefused
}
if errors.Is(usersErr, os.ErrNotExist) {
return nil
}
if usersErr != nil || removeAuthFile(usersPath) != nil {
return errCommandRefused
}
return nil
}
func setEnabled(directory string, enabled bool, args []string) error {
if len(args) != 1 || !usernamePattern.MatchString(args[0]) {
return errCommandRefused
}
return MutateUsers(directory, func(registry *Registry) error {
item := registry.FindByUsername(args[0])
if item == nil {
return errCommandRefused
}
if !enabled && item.Enabled && hasRole(item.Roles, RoleAdmin) && enabledAdminCount(*registry) == 1 {
return errLastAdministrator
}
item.Enabled = enabled
return nil
})
}
func changeRole(directory string, grant bool, args []string) error {
if len(args) != 3 || !usernamePattern.MatchString(args[0]) || args[1] != "--role" || (args[2] != string(RoleUser) && args[2] != string(RoleAdmin)) {
return errCommandRefused
}
role := Role(args[2])
return MutateUsers(directory, func(registry *Registry) error {
item := registry.FindByUsername(args[0])
if item == nil {
return errCommandRefused
}
if grant {
if hasRole(item.Roles, role) {
return errCommandRefused
}
item.Roles = append(item.Roles, role)
return nil
}
if !hasRole(item.Roles, role) || len(item.Roles) == 1 {
return errCommandRefused
}
if role == RoleAdmin && item.Enabled && enabledAdminCount(*registry) == 1 {
return errLastAdministrator
}
item.Roles = withoutRole(item.Roles, role)
return nil
})
}
func logoutAll(directory string, args []string) error {
if len(args) != 2 || !usernamePattern.MatchString(args[0]) || args[1] != "--yes" {
return errCommandRefused
}
return MutateUsers(directory, func(registry *Registry) error {
item := registry.FindByUsername(args[0])
if item == nil || item.AuthRevision == ^uint64(0) {
return errCommandRefused
}
item.AuthRevision++
return nil
})
}
func enabledAdminCount(registry Registry) int {
count := 0
for _, item := range registry.Users {
if item.Enabled && hasRole(item.Roles, RoleAdmin) {
count++
}
}
return count
}
func withoutRole(roles []Role, unwanted Role) []Role {
result := make([]Role, 0, len(roles)-1)
for _, role := range roles {
if role != unwanted {
result = append(result, role)
}
}
return result
}
func rolesToStrings(roles []Role) []string {
values := make([]string, len(roles))
for index, role := range roles {
values[index] = string(role)
}
return values
}
func validPublicURL(value string) bool {
parsed, err := url.Parse(value)
return err == nil && (parsed.Scheme == "http" || parsed.Scheme == "https") && parsed.Host != "" && parsed.User == nil && parsed.Path == "" && parsed.RawQuery == "" && parsed.Fragment == ""
}
func validOIDCURL(value string) bool {
parsed, err := url.Parse(value)
return err == nil && parsed.Scheme == "https" && parsed.Host != "" && parsed.User == nil && parsed.RawQuery == "" && parsed.Fragment == ""
}
func validHTTPSURL(value string) bool {
return validOIDCURL(value) && strings.TrimSuffix(value, "/") == value
}
func isHTTPS(value string) bool {
parsed, err := url.Parse(value)
return err == nil && parsed.Scheme == "https"
}
func isLoopbackURL(value string) bool {
parsed, err := url.Parse(value)
if err != nil {
return false
}
host := parsed.Hostname()
return host == "localhost" || net.ParseIP(host) != nil && net.ParseIP(host).IsLoopback()
}