Files
TA-Deployment-Broker/internal/app/authorization.go
T

645 lines
18 KiB
Go

package app
import (
"context"
"errors"
"fmt"
"net/http"
"strconv"
"strings"
"time"
)
func (s *Server) handlePortal(w http.ResponseWriter, r *http.Request) {
if hasAuditQuery(r) {
http.Redirect(w, r, "/portal/audit?"+r.URL.RawQuery, http.StatusSeeOther)
return
}
tech, err := s.currentTechnician(r)
if err != nil {
http.Redirect(w, r, "/", http.StatusSeeOther)
return
}
packages, err := s.listPackages(r)
if err != nil {
http.Error(w, "unable to list packages", http.StatusInternalServerError)
return
}
authorizations, err := s.listAuthorizations(r)
if err != nil {
http.Error(w, "unable to list authorizations", http.StatusInternalServerError)
return
}
actions, err := s.listActions(r)
if err != nil {
http.Error(w, "unable to list installer actions", http.StatusInternalServerError)
return
}
auditEvents, err := s.listAuditEvents(r, auditFilters{TimeRange: "all"}, 15)
if err != nil {
http.Error(w, "unable to list audit events", http.StatusInternalServerError)
return
}
s.render(w, "portal.html", pageData{
Title: "Deployment Access",
Technician: tech,
CSRFToken: tech.CSRFToken,
Authorizations: authorizations,
AuditEvents: auditEvents,
Packages: packages,
Actions: actions,
NewCode: r.URL.Query().Get("code"),
DefaultHostLimit: s.cfg.DefaultHostLimit,
DefaultDuration: strconv.Itoa(int(s.cfg.DefaultDuration.Hours())),
Notice: r.URL.Query().Get("notice"),
CurrentView: "codes",
})
}
func hasAuditQuery(r *http.Request) bool {
query := r.URL.Query()
for _, field := range []string{
"audit_range",
"audit_event",
"audit_customer",
"audit_user",
"audit_actor",
"audit_hostname",
"audit_package",
"audit_ip",
"audit_details",
} {
if query.Has(field) {
return true
}
}
return false
}
func (s *Server) handlePackagesPortal(w http.ResponseWriter, r *http.Request) {
tech, err := s.currentTechnician(r)
if err != nil {
http.Redirect(w, r, "/", http.StatusSeeOther)
return
}
packages, err := s.listPackages(r)
if err != nil {
http.Error(w, "unable to list packages", http.StatusInternalServerError)
return
}
s.render(w, "packages.html", pageData{
Title: "Protected Packages",
Technician: tech,
CSRFToken: tech.CSRFToken,
Packages: packages,
Notice: r.URL.Query().Get("notice"),
CurrentView: "packages",
})
}
func (s *Server) handleAuditPortal(w http.ResponseWriter, r *http.Request) {
tech, err := s.currentTechnician(r)
if err != nil {
http.Redirect(w, r, "/", http.StatusSeeOther)
return
}
auditFilter := auditFiltersFromRequest(r)
auditEventTypes, err := s.listAuditEventTypes(r)
if err != nil {
http.Error(w, "unable to list audit event types", http.StatusInternalServerError)
return
}
auditEvents, err := s.listAuditEvents(r, auditFilter, 250)
if err != nil {
http.Error(w, "unable to list audit events", http.StatusInternalServerError)
return
}
s.render(w, "audit.html", pageData{
Title: "Audit Trail",
Technician: tech,
CSRFToken: tech.CSRFToken,
AuditEvents: auditEvents,
AuditFilters: auditFilter,
AuditEventTypes: auditEventTypes,
CurrentView: "audit",
})
}
func (s *Server) listActions(r *http.Request) ([]actionRecord, error) {
rows, err := s.db.QueryContext(
r.Context(),
`SELECT slug, display_name, enabled
FROM installer_actions
ORDER BY sort_order, display_name`,
)
if err != nil {
return nil, err
}
defer rows.Close()
var records []actionRecord
for rows.Next() {
var record actionRecord
if err := rows.Scan(&record.Slug, &record.DisplayName, &record.Enabled); err != nil {
return nil, err
}
records = append(records, record)
}
return records, rows.Err()
}
func auditFiltersFromRequest(r *http.Request) auditFilters {
query := r.URL.Query()
timeRange := strings.TrimSpace(query.Get("audit_range"))
switch timeRange {
case "24h", "7d", "30d", "90d", "all":
default:
timeRange = "30d"
}
userFilter := query.Get("audit_user")
if strings.TrimSpace(userFilter) == "" {
// Preserve audit links created before the portal used "user" terminology.
userFilter = query.Get("audit_actor")
}
return auditFilters{
TimeRange: timeRange,
EventType: limitedFilter(query.Get("audit_event"), 100),
CustomerLabel: limitedFilter(query.Get("audit_customer"), 255),
User: limitedFilter(userFilter, 255),
Hostname: limitedFilter(query.Get("audit_hostname"), 255),
PackageSlug: limitedFilter(query.Get("audit_package"), 100),
SourceIP: limitedFilter(query.Get("audit_ip"), 64),
Details: limitedFilter(query.Get("audit_details"), 255),
}
}
func limitedFilter(value string, maxLength int) string {
value = strings.TrimSpace(value)
characters := []rune(value)
if len(characters) > maxLength {
return string(characters[:maxLength])
}
return value
}
func (s *Server) listAuditEventTypes(r *http.Request) ([]string, error) {
rows, err := s.db.QueryContext(
r.Context(),
`SELECT DISTINCT event_type
FROM audit_events
ORDER BY event_type`,
)
if err != nil {
return nil, err
}
defer rows.Close()
var eventTypes []string
for rows.Next() {
var eventType string
if err := rows.Scan(&eventType); err != nil {
return nil, err
}
eventTypes = append(eventTypes, eventType)
}
return eventTypes, rows.Err()
}
func auditQuery(filters auditFilters, limit int) (string, []any) {
query := `SELECT ae.event_type, COALESCE(a.customer_label, ''), ae.actor,
ae.hostname, ae.package_slug, ae.source_ip,
ae.details, ae.created_at
FROM audit_events ae
LEFT JOIN authorizations a ON a.id = ae.authorization_id
WHERE 1 = 1`
var arguments []any
timeClauses := map[string]string{
"24h": " AND ae.created_at >= datetime('now', '-1 day')",
"7d": " AND ae.created_at >= datetime('now', '-7 days')",
"30d": " AND ae.created_at >= datetime('now', '-30 days')",
"90d": " AND ae.created_at >= datetime('now', '-90 days')",
}
query += timeClauses[filters.TimeRange]
if filters.EventType != "" {
query += " AND ae.event_type = ?"
arguments = append(arguments, filters.EventType)
}
containsFilters := []struct {
column string
value string
}{
{"COALESCE(a.customer_label, '')", filters.CustomerLabel},
{"ae.actor", filters.User},
{"ae.hostname", filters.Hostname},
{"ae.package_slug", filters.PackageSlug},
{"ae.details", filters.Details},
}
for _, filter := range containsFilters {
if filter.value == "" {
continue
}
query += " AND instr(" + filter.column + ", ?) > 0"
arguments = append(arguments, filter.value)
}
if filters.SourceIP != "" {
query += " AND instr(ae.source_ip, ?) > 0"
arguments = append(arguments, filters.SourceIP)
}
query += " ORDER BY ae.created_at DESC LIMIT ?"
arguments = append(arguments, limit)
return query, arguments
}
func (s *Server) listAuditEvents(r *http.Request, filters auditFilters, limit int) ([]auditRecord, error) {
query, arguments := auditQuery(filters, limit)
rows, err := s.db.QueryContext(
r.Context(),
query,
arguments...,
)
if err != nil {
return nil, err
}
defer rows.Close()
var records []auditRecord
for rows.Next() {
var record auditRecord
if err := rows.Scan(
&record.EventType,
&record.CustomerLabel,
&record.User,
&record.Hostname,
&record.PackageSlug,
&record.SourceIP,
&record.Details,
&record.CreatedAt,
); err != nil {
return nil, err
}
records = append(records, record)
}
return records, rows.Err()
}
func (s *Server) listPackages(r *http.Request) ([]packageRecord, error) {
rows, err := s.db.QueryContext(
r.Context(),
`SELECT id, slug, display_name, package_name, package_version,
file_name, sha256, enabled
FROM packages
ORDER BY display_name, package_version`,
)
if err != nil {
return nil, err
}
defer rows.Close()
var records []packageRecord
for rows.Next() {
var record packageRecord
if err := rows.Scan(
&record.ID,
&record.Slug,
&record.DisplayName,
&record.PackageName,
&record.PackageVersion,
&record.FileName,
&record.SHA256,
&record.Enabled,
); err != nil {
return nil, err
}
records = append(records, record)
}
return records, rows.Err()
}
func (s *Server) listAuthorizations(r *http.Request) ([]authorizationRecord, error) {
rows, err := s.db.QueryContext(
r.Context(),
`SELECT a.id, a.code_hint, a.created_by, a.customer_label,
a.host_limit,
(SELECT COUNT(*) FROM authorization_hosts h
WHERE h.authorization_id = a.id),
a.created_at, a.expires_at, a.revoked_at,
COALESCE((
SELECT GROUP_CONCAT(display_name, ', ')
FROM (SELECT p.display_name
FROM authorization_packages ap
JOIN packages p ON p.id = ap.package_id
WHERE ap.authorization_id = a.id
ORDER BY p.display_name)
), ''),
COALESCE((
SELECT GROUP_CONCAT(display_name, ', ')
FROM (SELECT ia.display_name
FROM authorization_actions aa
JOIN installer_actions ia ON ia.slug = aa.action_slug
WHERE aa.authorization_id = a.id
ORDER BY ia.sort_order, ia.display_name)
), '')
FROM authorizations a
WHERE a.created_at > datetime('now', '-30 days')
ORDER BY a.created_at DESC
LIMIT 100`,
)
if err != nil {
return nil, err
}
defer rows.Close()
var records []authorizationRecord
for rows.Next() {
var record authorizationRecord
if err := rows.Scan(
&record.ID,
&record.CodeHint,
&record.CreatedBy,
&record.CustomerLabel,
&record.HostLimit,
&record.HostCount,
&record.CreatedAt,
&record.ExpiresAt,
&record.RevokedAt,
&record.Packages,
&record.Actions,
); err != nil {
return nil, err
}
records = append(records, record)
}
return records, rows.Err()
}
func (s *Server) handleCreateAuthorization(w http.ResponseWriter, r *http.Request) {
tech, _ := s.currentTechnician(r)
hostLimit, err := strconv.Atoi(r.FormValue("host_limit"))
if err != nil || hostLimit < 1 || hostLimit > s.cfg.MaxHostLimit {
http.Error(w, "invalid host limit", http.StatusBadRequest)
return
}
durationHours, err := strconv.Atoi(r.FormValue("duration_hours"))
if err != nil || durationHours < 1 || durationHours > 24 {
http.Error(w, "duration must be between 1 and 24 hours", http.StatusBadRequest)
return
}
packageIDs := r.Form["package_id"]
actionSlugs := r.Form["action_slug"]
if len(packageIDs) == 0 && len(actionSlugs) == 0 {
http.Error(w, "select at least one package or installer action", http.StatusBadRequest)
return
}
code, err := deploymentCode()
if err != nil {
http.Error(w, "unable to generate code", http.StatusInternalServerError)
return
}
codeHash := hashValue(code)
expiresAt := time.Now().UTC().Add(time.Duration(durationHours) * time.Hour)
customerLabel := strings.TrimSpace(r.FormValue("customer_label"))
if len(customerLabel) > 255 {
http.Error(w, "customer label is too long", http.StatusBadRequest)
return
}
tx, err := s.db.BeginTx(r.Context(), nil)
if err != nil {
http.Error(w, "unable to create authorization", http.StatusInternalServerError)
return
}
defer tx.Rollback()
result, err := tx.ExecContext(
r.Context(),
`INSERT INTO authorizations
(code_hash, code_hint, created_by, customer_label, host_limit, expires_at)
VALUES (?, ?, ?, ?, ?, ?)`,
codeHash[:], code[len(code)-5:], tech.Login, customerLabel, hostLimit, expiresAt,
)
if err != nil {
http.Error(w, "unable to create authorization", http.StatusInternalServerError)
return
}
authorizationID, _ := result.LastInsertId()
for _, rawID := range packageIDs {
packageID, err := strconv.ParseUint(rawID, 10, 64)
if err != nil {
http.Error(w, "invalid package selection", http.StatusBadRequest)
return
}
result, err := tx.ExecContext(
r.Context(),
`INSERT INTO authorization_packages
(authorization_id, package_id, package_slug, display_name,
package_name, package_version, file_name, sha256)
SELECT ?, id, slug, display_name, package_name, package_version,
file_name, sha256
FROM packages
WHERE id = ? AND enabled = TRUE`,
authorizationID, packageID,
)
if err != nil {
http.Error(w, "unable to authorize package", http.StatusInternalServerError)
return
}
affected, _ := result.RowsAffected()
if affected != 1 {
http.Error(w, "selected package is unavailable", http.StatusBadRequest)
return
}
}
for _, actionSlug := range actionSlugs {
actionSlug = strings.TrimSpace(actionSlug)
if !validSlug(actionSlug) {
http.Error(w, "invalid installer action", http.StatusBadRequest)
return
}
result, err := tx.ExecContext(
r.Context(),
`INSERT INTO authorization_actions (authorization_id, action_slug)
SELECT ?, slug FROM installer_actions
WHERE slug = ? AND enabled = TRUE`,
authorizationID, actionSlug,
)
if err != nil {
http.Error(w, "unable to authorize installer action", http.StatusInternalServerError)
return
}
affected, _ := result.RowsAffected()
if affected != 1 {
http.Error(w, "selected installer action is unavailable", http.StatusBadRequest)
return
}
}
if err := tx.Commit(); err != nil {
http.Error(w, "unable to save authorization", http.StatusInternalServerError)
return
}
id := uint64(authorizationID)
_ = s.audit(r.Context(), "authorization_created", tech.Login, &id, "", "", s.clientIP(r),
fmt.Sprintf("host_limit=%d duration_hours=%d customer=%q", hostLimit, durationHours, customerLabel))
http.Redirect(w, r, "/portal?code="+code, http.StatusSeeOther)
}
func (s *Server) handleRevokeAuthorization(w http.ResponseWriter, r *http.Request) {
tech, _ := s.currentTechnician(r)
id, err := parseUintPath(r, "id")
if err != nil {
http.Error(w, "invalid authorization", http.StatusBadRequest)
return
}
result, err := s.db.ExecContext(
r.Context(),
`UPDATE authorizations
SET revoked_at = COALESCE(revoked_at, CURRENT_TIMESTAMP)
WHERE id = ?`,
id,
)
if err != nil {
http.Error(w, "unable to revoke authorization", http.StatusInternalServerError)
return
}
affected, _ := result.RowsAffected()
if affected == 0 {
http.Error(w, "authorization not found", http.StatusNotFound)
return
}
_, _ = s.db.ExecContext(
r.Context(),
`UPDATE download_sessions SET revoked_at = CURRENT_TIMESTAMP
WHERE authorization_id = ? AND revoked_at IS NULL`,
id,
)
_ = s.audit(r.Context(), "authorization_revoked", tech.Login, &id, "", "", s.clientIP(r), "")
http.Redirect(w, r, "/portal?notice=Authorization+revoked", http.StatusSeeOther)
}
func (s *Server) handleUpsertPackage(w http.ResponseWriter, r *http.Request) {
tech, _ := s.currentTechnician(r)
fields := []string{
"slug", "display_name", "package_version", "file_name", "sha256",
}
values := make(map[string]string)
for _, field := range fields {
values[field] = strings.TrimSpace(r.FormValue(field))
if values[field] == "" {
http.Error(w, field+" is required", http.StatusBadRequest)
return
}
}
if !validSlug(values["slug"]) ||
len(values["display_name"]) > 255 ||
!validRegistrySegment(values["package_version"], 100) ||
!validRegistrySegment(values["file_name"], 255) ||
!validSHA256(values["sha256"]) {
http.Error(w, "invalid package metadata", http.StatusBadRequest)
return
}
enabled := r.FormValue("enabled") == "1"
err := s.savePackage(
r,
values["slug"],
values["display_name"],
values["slug"],
values["package_version"],
values["file_name"],
strings.ToLower(values["sha256"]),
enabled,
)
if err != nil {
http.Error(w, "unable to save package", http.StatusInternalServerError)
return
}
_ = s.audit(r.Context(), "package_saved", tech.Login, nil, "", values["slug"], s.clientIP(r),
fmt.Sprintf("version=%s enabled=%t", values["package_version"], enabled))
http.Redirect(w, r, "/portal/packages?notice=Package+saved", http.StatusSeeOther)
}
func (s *Server) savePackage(
r *http.Request,
slug string,
displayName string,
packageName string,
packageVersion string,
fileName string,
sha256 string,
enabled bool,
) error {
_, err := s.db.ExecContext(
r.Context(),
`INSERT INTO packages
(slug, display_name, package_name, package_version, file_name, sha256, enabled)
VALUES (?, ?, ?, ?, ?, ?, ?)
ON CONFLICT(slug) DO UPDATE SET
display_name = excluded.display_name,
package_name = excluded.package_name,
package_version = excluded.package_version,
file_name = excluded.file_name,
sha256 = excluded.sha256,
enabled = excluded.enabled,
updated_at = CURRENT_TIMESTAMP`,
slug,
displayName,
packageName,
packageVersion,
fileName,
sha256,
enabled,
)
return err
}
func validSlug(value string) bool {
if value == "" || len(value) > 100 {
return false
}
for _, character := range value {
if (character < 'a' || character > 'z') &&
(character < '0' || character > '9') &&
character != '-' && character != '_' {
return false
}
}
return true
}
func validSHA256(value string) bool {
if len(value) != 64 {
return false
}
for _, character := range strings.ToLower(value) {
if (character < '0' || character > '9') &&
(character < 'a' || character > 'f') {
return false
}
}
return true
}
func (s *Server) audit(
ctx context.Context,
eventType string,
user string,
authorizationID *uint64,
hostname string,
packageSlug string,
sourceIP string,
details string,
) error {
var authID any
if authorizationID != nil {
authID = *authorizationID
}
_, err := s.db.ExecContext(
ctx,
`INSERT INTO audit_events
(event_type, actor, authorization_id, hostname, package_slug,
source_ip, details)
VALUES (?, ?, ?, ?, ?, ?, ?)`,
eventType, user, authID, hostname, packageSlug, sourceIP, details,
)
return err
}
var errAuthorizationDenied = errors.New("authorization denied")