Files
cmbuyer/admin/cmd/device-credentials/main.go
T

201 lines
6.4 KiB
Go

package main
import (
"context"
"encoding/json"
"errors"
"flag"
"fmt"
"io"
"log"
"net/url"
"os"
"path/filepath"
"strings"
"time"
"cmbuyer/admin/internal/deviceauth"
"cmbuyer/admin/internal/storage/sqlite"
)
func main() {
if err := run(context.Background(), os.Args[1:], os.Stdout, os.Stderr); err != nil {
log.Print(err)
os.Exit(1)
}
}
func run(ctx context.Context, args []string, stdout, stderr io.Writer) error {
flags := flag.NewFlagSet("device-credentials", flag.ContinueOnError)
flags.SetOutput(stderr)
databaseSource := flags.String("database", "", "explicit migrated SQLite data source")
if err := flags.Parse(args); err != nil {
return err
}
if *databaseSource == "" {
return errors.New("-database is required")
}
if flags.NArg() < 1 {
return errors.New("usage: device-credentials -database <sqlite-data-source> <issue|list|revoke> [options]")
}
command := flags.Arg(0)
commandArgs := flags.Args()[1:]
var issueName, revokeDeviceID string
switch command {
case "issue":
commandFlags := flag.NewFlagSet("issue", flag.ContinueOnError)
commandFlags.SetOutput(stderr)
commandFlags.StringVar(&issueName, "name", "", "non-secret device display name")
if err := commandFlags.Parse(commandArgs); err != nil {
return err
}
if issueName == "" || commandFlags.NArg() != 0 {
return errors.New("usage: device-credentials -database <sqlite-data-source> issue -name <display-name>")
}
case "list":
if len(commandArgs) != 0 {
return errors.New("usage: device-credentials -database <sqlite-data-source> list")
}
case "revoke":
commandFlags := flag.NewFlagSet("revoke", flag.ContinueOnError)
commandFlags.SetOutput(stderr)
commandFlags.StringVar(&revokeDeviceID, "device-id", "", "canonical device UUID")
if err := commandFlags.Parse(commandArgs); err != nil {
return err
}
if revokeDeviceID == "" || commandFlags.NArg() != 0 {
return errors.New("usage: device-credentials -database <sqlite-data-source> revoke -device-id <uuid>")
}
default:
return fmt.Errorf("unsupported device credential command %q", command)
}
if command == "issue" && !deviceauth.ValidDisplayName(issueName) {
return deviceauth.ErrInvalidCredential
}
if command == "revoke" && !deviceauth.ValidDeviceID(revokeDeviceID) {
return deviceauth.ErrInvalidCredential
}
existingSource, err := existingSQLiteDataSource(*databaseSource)
if err != nil {
return err
}
database, err := sqlite.Open(existingSource)
if err != nil {
return fmt.Errorf("open SQLite database: %w", err)
}
defer database.Close()
store, err := deviceauth.NewCredentialStore(database)
if err != nil {
return err
}
switch command {
case "issue":
issued, err := store.Issue(ctx, issueName)
if err != nil {
return err
}
// The token has json:"-" and is printed only by this explicit post-commit path. Generic
// serialization, list, revoke, errors, and server responses therefore cannot disclose it.
if _, err := fmt.Fprintf(stdout, "device_id=%s\ndisplay_name=%s\ntoken=%s\ncreated_at=%s\n",
issued.DeviceID, issued.DisplayName, issued.Token, issued.CreatedAt.Format(time.RFC3339Nano)); err != nil {
return errors.New("write issued device credential")
}
return nil
case "list":
credentials, err := store.List(ctx)
if err != nil {
return err
}
return writeJSON(stdout, credentials)
case "revoke":
credential, changed, err := store.Revoke(ctx, revokeDeviceID)
if err != nil {
return err
}
return writeJSON(stdout, struct {
Credential deviceauth.Credential `json:"credential"`
RevokedNow bool `json:"revoked_now"`
}{Credential: credential, RevokedNow: changed})
}
return errors.New("unreachable device credential command")
}
func existingSQLiteDataSource(value string) (string, error) {
if value == "" || strings.TrimSpace(value) != value {
return "", errors.New("-database must name an existing file-backed SQLite database")
}
var parsed *url.URL
var query url.Values
if strings.HasPrefix(strings.ToLower(value), "file:") {
var err error
parsed, err = url.Parse(value)
if err != nil || !strings.EqualFold(parsed.Scheme, "file") || parsed.User != nil || parsed.Host != "" || parsed.Fragment != "" {
return "", errors.New("-database file URI is invalid")
}
// go-sqlite3 recognizes URI filenames only with the exact lowercase file: prefix.
// Canonicalize accepted scheme casing before mode=rw reaches the driver, otherwise a
// mixed-case input could be treated as a plain filename and recreate a missing database.
parsed.Scheme = "file"
query, err = url.ParseQuery(parsed.RawQuery)
if err != nil {
return "", errors.New("-database query parameters are invalid")
}
fileName := parsed.Path
if parsed.Opaque != "" {
fileName = parsed.Opaque
}
decodedName, err := url.PathUnescape(fileName)
if err != nil || fileName == "" || strings.EqualFold(decodedName, ":memory:") {
return "", errors.New("-database must name an existing file-backed SQLite database")
}
} else {
pathPart, rawQuery, hasQuery := strings.Cut(value, "?")
if pathPart == "" || strings.EqualFold(pathPart, ":memory:") || strings.Contains(pathPart, "://") {
return "", errors.New("-database must name an existing file-backed SQLite database")
}
var err error
query, err = url.ParseQuery(rawQuery)
if err != nil {
return "", errors.New("-database query parameters are invalid")
}
normalizedPath := filepath.ToSlash(pathPart)
if filepath.VolumeName(pathPart) != "" && !strings.HasPrefix(normalizedPath, "/") {
normalizedPath = "/" + normalizedPath
}
parsed = &url.URL{Scheme: "file", Path: normalizedPath}
if !hasQuery {
query = make(url.Values)
}
}
modes := query["mode"]
if len(modes) > 1 || len(modes) == 1 && modes[0] != "rw" {
return "", errors.New("-database only permits SQLite mode=rw")
}
if len(modes) == 0 {
query.Set("mode", "rw")
}
for _, name := range []string{"immutable", "_query_only"} {
for _, setting := range query[name] {
if setting != "0" && !strings.EqualFold(setting, "false") {
return "", errors.New("-database contains a read-only SQLite option")
}
}
}
parsed.RawQuery = query.Encode()
parsed.ForceQuery = false
return parsed.String(), nil
}
func writeJSON(writer io.Writer, value any) error {
encoder := json.NewEncoder(writer)
encoder.SetEscapeHTML(true)
if err := encoder.Encode(value); err != nil {
return errors.New("write device credential metadata")
}
return nil
}